diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 647e04936..ded77ec58 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -46,3 +46,25 @@ updates: wireguard: patterns: - "golang.zx2c4.com/wireguard*" + + # Base images of the source-build Dockerfiles, pinned by digest (Chainguard + # publishes only :latest for free). Dockerfile.release files feed goreleaser + # and keep the published images as they are, so their bases are left alone. + - package-ecosystem: "docker" + directories: + - "/upload-server" + schedule: + interval: "weekly" + open-pull-requests-limit: 3 + groups: + base-images: + patterns: + - "*" + ignore: + - dependency-name: "gcr.io/distroless/base" + # Go minor and major versions move with the rest of the repository; + # patch releases and new digests of the pinned tag still come through. + - dependency-name: "golang" + update-types: + - "version-update:semver-minor" + - "version-update:semver-major" diff --git a/.github/scripts/test-homebrew-cask.sh b/.github/scripts/test-homebrew-cask.sh new file mode 100755 index 000000000..9c1041bdd --- /dev/null +++ b/.github/scripts/test-homebrew-cask.sh @@ -0,0 +1,338 @@ +#!/usr/bin/env bash +set -euo pipefail + +fail() { + echo "::error::$*" >&2 + exit 1 +} + +if [[ ${RUNNER_ENVIRONMENT:-} != github-hosted || ${RUNNER_OS:-} != macOS || $(uname -s) != Darwin ]]; then + fail "This test installs a system daemon and must run on a disposable GitHub macOS runner." +fi +if [[ $EUID == 0 ]]; then + fail "Run this script as the Homebrew user, not root." +fi + +readonly test_dir="${RUNNER_TEMP:?}/homebrew-cask" +readonly results_dir="$test_dir/results" +readonly app='/Applications/Netbird UI.app' +readonly plist='/Library/LaunchDaemons/netbird.plist' +readonly cask='netbirdio/tap/netbird-ui' +readonly formula='netbirdio/tap/netbird' +readonly published_cask="$test_dir/published-netbird-ui.rb" +readonly legacy_cask="$test_dir/legacy-netbird-ui.rb" +readonly rendered_cask="$test_dir/rendered-netbird-ui.rb" +readonly fixture_dir="$test_dir/fixture" +readonly serve_dir="$test_dir/serve" +readonly fixture_zip="$serve_dir/netbird-ui.zip" +readonly fixture_port=18080 +readonly fixture_url="http://127.0.0.1:$fixture_port/netbird-ui.zip" +readonly marker="$test_dir/installer.marker" + +mkdir -p "$results_dir" "$fixture_dir/netbird_ui_darwin" "$serve_dir" "$test_dir/downloads" +exec > >(tee "$results_dir/test.log") 2>&1 + +sudo -n true +if command -v netbird || [[ -e "$app" || -e "$plist" ]] || pgrep -x netbird-ui; then + fail "The runner already has NetBird installed or running." +fi +if sudo launchctl print system/netbird > "$results_dir/initial-service.log" 2>&1; then + fail "The runner already has a NetBird service loaded." +fi + +install_attempted=false +server_pid='' +daemon_pid='' +version='' + +stop_ui() { + local status=0 + sudo pkill -x netbird-ui || status=$? + # pkill returns 1 when the UI is already closed. + [[ $status == 0 || $status == 1 ]] +} + +cleanup() { + local status=$? + trap - EXIT + set +e + + if [[ $install_attempted == true ]]; then + stop_ui || status=1 + if [[ -S /var/run/netbird.sock ]]; then + sudo netbird down || status=1 + fi + if brew list --cask "$cask" >/dev/null 2>&1 || [[ -e "$app" ]]; then + brew uninstall --cask --force "$cask" || status=1 + fi + # A failed cask install can leave a daemon even after Homebrew rolls back the app. + if sudo launchctl print system/netbird > "$results_dir/cleanup-service.log" 2>&1; then + sudo netbird service stop || status=1 + fi + if [[ -e "$plist" ]]; then + sudo netbird service uninstall || status=1 + fi + fi + if [[ -f /var/log/netbird/client.log ]]; then + sudo cat /var/log/netbird/client.log > "$results_dir/client.log" || status=1 + fi + if command -v netbird >/dev/null; then + brew uninstall --formula "$formula" || status=1 + fi + if [[ -n $server_pid ]]; then + kill "$server_pid" 2>/dev/null || true + fi + exit "$status" +} +trap cleanup EXIT +trap 'exit 130' INT +trap 'exit 143' TERM + +run_logged() { + local name=$1 + shift + "$@" 2>&1 | tee "$results_dir/$name.log" +} + +cask_field() { + local stanza=$1 file=$2 + sed -nE "s/^[[:space:]]*$stanza \"([^\"]+)\".*/\\1/p" "$file" +} + +release_fields() { + local file=$1 + grep -E '^[[:space:]]*(version|url|sha256|app) ' "$file" +} + +use_cask() { + local file=$1 + cp "$file" "$tap_dir/Casks/netbird-ui.rb" +} + +# The released installer opens the UI as root, which never returns on a headless +# runner. The cask only needs two script paths and a version argument, so the test +# ships a stub bundle that records what it received and starts the daemon. +build_fixture() { + local bundle="$fixture_dir/netbird_ui_darwin" + printf '#!/bin/sh\nexit 0\n' > "$bundle/netbird-ui" + chmod 755 "$bundle/netbird-ui" + # After a bootout launchd keeps tearing the previous daemon down for a couple of + # seconds, and loading the same label again fails until that finishes. + cat > "$bundle/installer.sh" < '$marker' +netbird service install +attempt=0 +until netbird service start; do + attempt=\$((attempt + 1)) + [ "\$attempt" -lt 15 ] || exit 1 + sleep 1 +done +EOF + printf '#!/bin/sh\nexit 0\n' > "$bundle/uninstaller.sh" + # Shipped without the executable bit so the 0755 seen after install can only come from the cask. + chmod 644 "$bundle/installer.sh" "$bundle/uninstaller.sh" + rm -f "$fixture_zip" + (cd "$fixture_dir" && zip -qr "$fixture_zip" netbird_ui_darwin) +} + +start_fixture_server() { + python3 -m http.server "$fixture_port" --bind 127.0.0.1 --directory "$serve_dir" \ + > "$results_dir/fixture-server.log" 2>&1 & + server_pid=$! + local attempt + for attempt in {1..20}; do + if curl --silent --fail --output /dev/null "$fixture_url"; then + return + fi + sleep 0.5 + done + fail "The fixture HTTP server did not come up on port $fixture_port." +} + +assert_published_layout() { + local url archive script + while read -r url; do + archive="$test_dir/downloads/${url##*/}" + curl --fail --location --silent --retry 3 --output "$archive" "$url" + for script in installer.sh uninstaller.sh; do + unzip -l "$archive" | grep -q " netbird_ui_darwin/$script\$" || + fail "The published archive ${url##*/} has no netbird_ui_darwin/$script." + done + done < <(cask_field url "$published_cask") +} + +assert_no_deprecations() { + if grep -Ei '(postflight|uninstall_preflight).*deprecated|deprecated.*(postflight|uninstall_preflight)' "$@"; then + fail "Homebrew reported a deprecated cask lifecycle hook." + fi +} + +wait_for_daemon() { + local attempt + for attempt in {1..30}; do + if sudo launchctl print system/netbird > "$results_dir/service.log" 2>&1 && + grep -Eq '^[[:space:]]*state = running$' "$results_dir/service.log"; then + return + fi + sleep 1 + done + cat "$results_dir/service.log" + fail "The installed daemon did not reach the running state." +} + +wait_for_exit() { + local pid=$1 attempt + for attempt in {1..30}; do + if ! sudo kill -0 "$pid" 2>/dev/null; then + return + fi + sleep 1 + done + fail "Daemon process $pid is still running after removal." +} + +assert_service_absent() { + if sudo launchctl print system/netbird > "$results_dir/removed-service.log" 2>&1; then + fail "The NetBird service is still loaded after removal." + fi +} + +assert_installed() { + local script + [[ -f $marker ]] || fail "The cask did not run installer.sh." + grep -qx "version=$version" "$marker" || fail "installer.sh did not receive the cask version: $(cat "$marker")" + grep -qx 'uid=0' "$marker" || fail "installer.sh did not run as root: $(cat "$marker")" + [[ -d "$app" && -x "$app/netbird-ui" ]] || fail "The UI was not installed." + for script in installer.sh uninstaller.sh; do + [[ $(stat -f '%Lp' "$app/$script") == 755 ]] || fail "Incorrect permissions on $script." + done + [[ -f "$plist" ]] || fail "The installer did not create the daemon plist." + wait_for_daemon + daemon_pid=$(awk '/^[[:space:]]*pid = / { print $3; exit }' "$results_dir/service.log") + [[ $daemon_pid =~ ^[0-9]+$ ]] || fail "The running daemon has no PID." + sudo kill -0 "$daemon_pid" +} + +assert_uninstalled() { + local log=$1 + assert_no_deprecations "$log" + [[ ! -e "$app" ]] || fail "The UI app remains after uninstall." + [[ ! -e "$plist" ]] || fail "The daemon plist remains after uninstall." + assert_service_absent + wait_for_exit "$daemon_pid" + [[ $(netbird version) == "$version" ]] || fail "Cask uninstall removed the CLI dependency." +} + +installed_caskfiles() { + local extension=$1 + find "$(brew --caskroom)/netbird-ui/.metadata" -name "netbird-ui.$extension" 2>/dev/null +} + +assert_legacy_metadata() { + installed_caskfiles rb | grep -q . || fail "The legacy cask did not leave a Ruby caskfile behind." +} + +assert_steps_metadata() { + if installed_caskfiles rb | grep -q .; then + fail "Homebrew still keeps the legacy Ruby caskfile after reinstall." + fi + installed_caskfiles json | grep -q . || fail "Homebrew did not save the reinstalled cask as JSON." +} + +brew --version +sw_vers +brew tap netbirdio/tap "${GITHUB_WORKSPACE:?}/.homebrew-cask-tap" +tap_dir=$(brew --repository netbirdio/tap) +readonly tap_dir + +[[ -f "$tap_dir/Casks/netbird-ui.rb" ]] || fail "The tap has no Casks/netbird-ui.rb." +cp "$tap_dir/Casks/netbird-ui.rb" "$published_cask" +cp "$published_cask" "$results_dir/published-netbird-ui.rb" + +version=$(brew info --json=v2 --formula "$formula" | jq -r '.formulae[0].versions.stable') +readonly version +[[ -n $version && $version != null ]] || fail "Could not read the formula version from the tap." + +assert_published_layout + +build_fixture +fixture_sha=$(shasum -a 256 "$fixture_zip" | cut -d' ' -f1) +readonly fixture_sha +start_fixture_server + +export PROJECT=netbird-ui VERSION="$version" +export AMD="$fixture_zip" ARM="$fixture_zip" AMD_URL="$fixture_url" ARM_URL="$fixture_url" +gomplate -f "$GITHUB_WORKSPACE/client/ui/netbird-ui.rb.tmpl" -o "$rendered_cask" +cp "$rendered_cask" "$results_dir/rendered-netbird-ui.rb" + +sed -E "s|^([[:space:]]*version) \"[^\"]+\"|\\1 \"$version\"|; s|^([[:space:]]*url) \"[^\"]+\"|\\1 \"$fixture_url\"|; s|^([[:space:]]*sha256) \"[^\"]+\"|\\1 \"$fixture_sha\"|" \ + "$published_cask" > "$legacy_cask" +cp "$legacy_cask" "$results_dir/legacy-netbird-ui.rb" +if ! diff <(release_fields "$legacy_cask") <(release_fields "$rendered_cask"); then + fail "The rendered cask changes release data, not only lifecycle stanzas." +fi + +use_cask "$rendered_cask" +brew info --json=v2 --cask "$cask" > "$results_dir/cask.json" 2> "$results_dir/load.log" +cat "$results_dir/load.log" +assert_no_deprecations "$results_dir/load.log" +run_logged style brew style --cask --only-cops=Cask/InstallSteps "$cask" + +run_logged install-cli brew install --formula "$formula" +[[ $(netbird version) == "$version" ]] || fail "The installed CLI does not report the formula version." + +for scenario in running stopped missing; do + echo "::group::Uninstall with $scenario service" + install_attempted=true + sudo rm -f "$marker" + run_logged "install-$scenario" brew install --cask "$cask" + assert_no_deprecations "$results_dir/install-$scenario.log" + assert_installed + stop_ui + + case "$scenario" in + running) ;; + stopped) + run_logged stop-daemon sudo netbird service stop + wait_for_exit "$daemon_pid" + [[ -f "$plist" ]] || fail "Stopping the daemon unexpectedly removed its plist." + ;; + missing) + run_logged stop-missing-daemon sudo netbird service stop + run_logged remove-daemon sudo netbird service uninstall + wait_for_exit "$daemon_pid" + [[ ! -e "$plist" ]] || fail "The missing-service scenario still has a plist." + assert_service_absent + ;; + *) fail "Unknown uninstall scenario: $scenario" ;; + esac + + run_logged "uninstall-$scenario" brew uninstall --cask "$cask" + assert_uninstalled "$results_dir/uninstall-$scenario.log" + echo "::endgroup::" +done + +# Every existing user first meets the new cask through an upgrade of the published +# one, whose legacy flight blocks Homebrew replays from the saved Ruby caskfile. +echo "::group::Reinstall over the published legacy cask" +install_attempted=true +use_cask "$legacy_cask" +sudo rm -f "$marker" +run_logged install-legacy brew install --cask "$cask" +assert_installed +assert_legacy_metadata +stop_ui + +use_cask "$rendered_cask" +sudo rm -f "$marker" +run_logged reinstall-legacy brew reinstall --cask "$cask" +assert_installed +assert_steps_metadata +stop_ui + +run_logged uninstall-legacy brew uninstall --cask "$cask" +assert_uninstalled "$results_dir/uninstall-legacy.log" +echo "::endgroup::" diff --git a/.github/workflows/frontend-ui.yml b/.github/workflows/frontend-ui.yml index 014c5c2ae..2ad43c581 100644 --- a/.github/workflows/frontend-ui.yml +++ b/.github/workflows/frontend-ui.yml @@ -38,12 +38,12 @@ jobs: persist-credentials: false - name: Set up Node.js - uses: actions/setup-node@v4 + uses: actions/setup-node@v7 with: node-version: "22" - name: Set up pnpm - uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 + uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0 with: version: 11 @@ -79,7 +79,7 @@ jobs: run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT" - name: Cache pnpm store - uses: actions/cache@v4 + uses: actions/cache@v6 with: path: ${{ steps.pnpm-store.outputs.path }} key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }} diff --git a/.github/workflows/golang-test-freebsd.yml b/.github/workflows/golang-test-freebsd.yml index 65c39147a..7bd48e3d0 100644 --- a/.github/workflows/golang-test-freebsd.yml +++ b/.github/workflows/golang-test-freebsd.yml @@ -33,7 +33,7 @@ jobs: with: usesh: true copyback: false - release: "15.0" + release: "15.1" envs: "GO_VERSION" prepare: | pkg install -y curl pkgconf xorg diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 586e1235b..ff36a0854 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -80,3 +80,49 @@ jobs: skip-save-cache: true cache-invalidation-interval: 0 args: --timeout=20m + + # Separate job rather than extra rows in the matrix above: those rows pick a + # GOOS by picking a runner OS, while android/ios are cross-compiled from + # ubuntu — an `include` entry with os: ubuntu-latest would merge into the + # Linux row instead of adding one. The package path is restricted because a + # whole-repo run under GOOS=android pulls *_linux.go files into packages that + # have no android counterpart. + golangci-mobile: + strategy: + fail-fast: false + matrix: + include: + - goos: android + goarch: arm64 + packages: ./client/android/... + display_name: Android + - goos: ios + goarch: arm64 + packages: ./client/ios/... + display_name: iOS + name: ${{ matrix.display_name }} + runs-on: ubuntu-latest + timeout-minutes: 25 + env: + CGO_ENABLED: 0 + GOOS: ${{ matrix.goos }} + GOARCH: ${{ matrix.goarch }} + steps: + - name: Checkout code + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + persist-credentials: false + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version-file: "go.mod" + cache: false + - name: golangci-lint + uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1 + with: + version: latest + install-mode: binary + skip-cache: true + skip-save-cache: true + cache-invalidation-interval: 0 + args: --timeout=20m ${{ matrix.packages }} diff --git a/.github/workflows/mobile-build-validation.yml b/.github/workflows/mobile-build-validation.yml new file mode 100644 index 000000000..613a39b3d --- /dev/null +++ b/.github/workflows/mobile-build-validation.yml @@ -0,0 +1,64 @@ +name: Mobile + +on: + push: + branches: + - main + - "release-*" + pull_request: + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }} + cancel-in-progress: true + +jobs: + android_build: + name: "Android / Build" + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + goarch: [arm64, arm, amd64, "386"] + env: + CGO_ENABLED: 0 + GOOS: android + GOARCH: ${{ matrix.goarch }} + steps: + - name: Checkout repository + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + persist-credentials: false + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version-file: "go.mod" + - name: Build Android bridge + run: go build ./client/android/... + - name: Vet Android bridge + if: matrix.goarch == 'arm64' + run: go vet ./client/android/... + + ios_build: + name: "iOS / Build" + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + goarch: [arm64, amd64] + env: + CGO_ENABLED: 0 + GOOS: ios + GOARCH: ${{ matrix.goarch }} + steps: + - name: Checkout repository + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + persist-credentials: false + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version-file: "go.mod" + # No `go vet` counterpart: every ios target requires external (cgo) + # linking, which needs an Xcode toolchain the runner does not have. + - name: Build iOS SDK + run: go build ./client/ios/... diff --git a/.github/workflows/redhat-certify.yml b/.github/workflows/redhat-certify.yml new file mode 100644 index 000000000..e592dabc2 --- /dev/null +++ b/.github/workflows/redhat-certify.yml @@ -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_ +# 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 diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 48f439be9..2d2619ecd 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -69,7 +69,7 @@ jobs: with: usesh: true copyback: false - release: "15.0" + release: "15.1" envs: "GO_VERSION" prepare: | # Install required packages @@ -191,6 +191,17 @@ jobs: # 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 uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0 - name: Set up Docker Buildx @@ -230,14 +241,18 @@ jobs: uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2 with: version: ${{ env.GORELEASER_VER }} - args: release --clean ${{ env.flags }} + args: release --config .goreleaser.generated.yaml --clean ${{ env.flags }} env: GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }} UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }} UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }} GPG_RPM_KEY_FILE: ${{ env.GPG_RPM_KEY_FILE }} - NFPM_NETBIRD_RPM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} + # One per nfpm id: GoReleaser looks the passphrase up as NFPM__PASSPHRASE. + NFPM_NETBIRD_RPM_AMD64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} + NFPM_NETBIRD_RPM_ARM64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} + NFPM_NETBIRD_RPM_ARM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} + NFPM_NETBIRD_RPM_386_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }} SKIP_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }} - name: Verify RPM signatures @@ -294,10 +309,12 @@ jobs: tag_and_push() { local src="$1" img_name tag dst variant="" img_name="${src%%:*}" - # Client variants share a repository, so keep their tag suffixes. + # Variants share a repository with their default image, so keep + # their tag suffixes. Order matters: the first matching pattern wins. case "$src" in *-rootless-ubi-amd64) variant="-rootless-ubi" ;; *-rootless-amd64) variant="-rootless" ;; + *-ubi-amd64) variant="-ubi" ;; esac for tag in $(resolve_tags); do dst="${img_name}:${tag}${variant}" @@ -365,6 +382,24 @@ jobs: path: dist/netbird_darwin** retention-days: 7 + # Certify the UBI images in the Red Hat Ecosystem Catalog on stable tags. + # See redhat-certify.yml, which can also be run by hand for any released version. + redhat_certification: + name: "Red Hat" + needs: release + if: | + github.repository == 'netbirdio/netbird' && + startsWith(github.ref, 'refs/tags/v') && + !contains(github.ref_name, '-') + permissions: + contents: read + uses: ./.github/workflows/redhat-certify.yml + with: + component: all + version: ${{ github.ref_name }} + secrets: + PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }} + release_ui: runs-on: ubuntu-latest outputs: @@ -419,12 +454,12 @@ jobs: run: git --no-pager diff --exit-code - name: Set up Node.js - uses: actions/setup-node@v4 + uses: actions/setup-node@v7 with: node-version: '22' - name: Set up pnpm - uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 + uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0 with: version: 11 @@ -556,12 +591,12 @@ jobs: run: git --no-pager diff --exit-code - name: Set up Node.js - uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0 with: node-version: '22' - name: Set up pnpm - uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 + uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0 with: version: 11 @@ -653,11 +688,11 @@ jobs: - name: check git status run: git --no-pager diff --exit-code - name: Set up Node.js - uses: actions/setup-node@v4 + uses: actions/setup-node@v7 with: node-version: '22' - name: Set up pnpm - uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 + uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0 with: version: 11 - name: Install wails3 CLI @@ -776,7 +811,7 @@ jobs: run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z" - name: Set up Go for wails3 CLI - uses: actions/setup-go@v5 + uses: actions/setup-go@v6 with: go-version-file: "go.mod" cache: false diff --git a/.github/workflows/test-homebrew-cask.yml b/.github/workflows/test-homebrew-cask.yml new file mode 100644 index 000000000..a75951425 --- /dev/null +++ b/.github/workflows/test-homebrew-cask.yml @@ -0,0 +1,46 @@ +name: Test Homebrew cask + +on: + pull_request: + paths: + - "client/ui/netbird-ui.rb.tmpl" + - ".github/scripts/test-homebrew-cask.sh" + - ".github/workflows/test-homebrew-cask.yml" + workflow_dispatch: + +permissions: + contents: read + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }} + cancel-in-progress: true + +jobs: + install-uninstall: + runs-on: macos-latest + timeout-minutes: 20 + steps: + - name: Checkout code + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + persist-credentials: false + + - name: Clone the Homebrew tap + run: git clone https://github.com/netbirdio/homebrew-tap.git .homebrew-cask-tap + + - name: Update Homebrew and install gomplate + # The runner image disables auto-update; the cask steps DSL needs Homebrew 6.0.20 or newer. + run: | + brew update + brew install gomplate + + - name: Install and uninstall the cask + run: .github/scripts/test-homebrew-cask.sh + + - name: Upload logs + if: always() + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a #v7.0.1 + with: + name: homebrew-cask-results + path: ${{ runner.temp }}/homebrew-cask/results + if-no-files-found: ignore diff --git a/.github/workflows/ui-translations.yml b/.github/workflows/ui-translations.yml index 7d3b12f2d..24b7c9de2 100644 --- a/.github/workflows/ui-translations.yml +++ b/.github/workflows/ui-translations.yml @@ -32,7 +32,7 @@ jobs: persist-credentials: false - name: Set up Node.js - uses: actions/setup-node@v4 + uses: actions/setup-node@v7 with: node-version: "22" diff --git a/.gitignore b/.gitignore index dd7eea76f..5c01f6e60 100644 --- a/.gitignore +++ b/.gitignore @@ -38,4 +38,7 @@ management/server/types/testdata/ # generated by chglog in the release workflow, embedded into the RPM changelog.yml + +# generated by rpm-provides.sh, the config GoReleaser actually runs +.goreleaser.generated.yaml .chglog.yml diff --git a/.goreleaser.yaml b/.goreleaser.yaml index 4a7141866..eafb09c3e 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -60,6 +60,34 @@ builds: - load_wgnt_from_rsrc - pkcs11 + # Single-arch builds: nfpm provides is not templated, so the RPM splits per arch. They + # carry the PKCS#11 store like the deb build above. + - &netbird_rpm_build + id: netbird-rpm-amd64 + dir: client + binary: netbird + env: [CGO_ENABLED=0] + goos: [linux] + goarch: [amd64] + ldflags: + - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser + mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - load_wgnt_from_rsrc + - pkcs11 + + - <<: *netbird_rpm_build + id: netbird-rpm-arm64 + goarch: [arm64] + + - <<: *netbird_rpm_build + id: netbird-rpm-arm + goarch: [arm] + + - <<: *netbird_rpm_build + id: netbird-rpm-386 + goarch: [386] + - id: netbird-static dir: client binary: netbird @@ -243,17 +271,22 @@ nfpms: postinstall: "release_files/post_install.sh" preremove: "release_files/pre_remove.sh" - - maintainer: Netbird + - &netbird_rpm + maintainer: Netbird description: Netbird client. homepage: https://netbird.io/ license: BSD-3-Clause vendor: NetBird - id: netbird_rpm + id: netbird_rpm_amd64 bindir: /usr/bin - builds: - - netbird-pkcs11 + ids: + - netbird-rpm-amd64 formats: - rpm + # Red Hat certification (RPM Version Handling) requires rpmbuild's ISA + # provide, which nfpm does not emit. The version is filled in by the release job. + provides: + - "netbird(x86-64) = @RPM_EVR@" # The client verifies TLS to management and signal against the system trust # store. Red Hat software certification (RPM Dependency Tracking) also # rejects packages that declare no dependencies at all. @@ -283,6 +316,27 @@ nfpms: packager: NetBird signature: key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}' + + - <<: *netbird_rpm + id: netbird_rpm_arm64 + ids: + - netbird-rpm-arm64 + provides: + - "netbird(aarch-64) = @RPM_EVR@" + + - <<: *netbird_rpm + id: netbird_rpm_arm + ids: + - netbird-rpm-arm + provides: + - "netbird(armv6hl-32) = @RPM_EVR@" + + - <<: *netbird_rpm + id: netbird_rpm_386 + ids: + - netbird-rpm-386 + provides: + - "netbird(x86-32) = @RPM_EVR@" dockers_v2: - id: netbird disable: "{{ .Env.SKIP_DOCKER_PUSH }}" @@ -445,7 +499,7 @@ dockers_v2: tags: - "{{ .Version }}" - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}" - dockerfile: upload-server/Dockerfile + dockerfile: upload-server/Dockerfile.release platforms: - linux/amd64 - linux/arm64 @@ -501,6 +555,41 @@ dockers_v2: "org.opencontainers.image.revision": "{{.FullCommit}}" "org.opencontainers.image.source": "{{.GitURL}}" "maintainer": "dev@netbird.io" + - id: proxy-ubi + disable: "{{ .Env.SKIP_DOCKER_PUSH }}" + ids: + - netbird-proxy + images: + - netbirdio/reverse-proxy + - ghcr.io/netbirdio/reverse-proxy + tags: + - "{{ .Version }}-ubi" + - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}" + dockerfile: proxy/Dockerfile.ubi + platforms: + - linux/amd64 + - linux/arm64 + build_args: + VERSION: "{{ .Version }}" + RELEASE: "{{ .Timestamp }}" + hooks: + pre: + - cmd: 'sh 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: - ids: @@ -533,7 +622,10 @@ uploads: - name: yum skip: "{{ .Env.SKIP_PUBLISH }}" ids: - - netbird_rpm + - netbird_rpm_amd64 + - netbird_rpm_arm64 + - netbird_rpm_arm + - netbird_rpm_386 mode: archive target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }} username: dev@wiretrustee.com diff --git a/README.md b/README.md index 336332043..4dbfc7bd0 100644 --- a/README.md +++ b/README.md @@ -115,6 +115,56 @@ export NETBIRD_DOMAIN=netbird.example.com; curl -fsSL https://github.com/netbird See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details. +### Reporting bugs and requesting features + +NetBird uses a discussion-first workflow. Bug reports and feature requests start in +[Discussions](https://github.com/netbirdio/netbird/discussions), not as issues. + +| What you want to do | Where to go | +| --- | --- | +| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) | +| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) | +| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) | +| Report a security vulnerability | [Security policy](https://github.com/netbirdio/netbird/security/policy), never a public thread | + +Our team and maintainers triage discussions, ask follow-up questions, check for duplicates, +and reproduce bugs. Validated reports are promoted to issues. This keeps the issue tracker a clear +answer to one question: what is the team working on. + +Please search existing discussions and issues first, including closed ones. If something similar +already exists, upvote it and add your details there instead of opening a duplicate. + +For bug reports, include your NetBird version, operating system, deployment type (Cloud, +self-hosted, Kubernetes, or Docker), reproduction steps, expected and actual behavior, and a debug +bundle where relevant: + +```shell +netbird version +netbird status -d -A +netbird debug for 1m -A -S -U +``` + +`-U` uploads the bundle and prints a file key you can paste instead of attaching the archive. +`-A` anonymizes the output, which matters on a public thread. It masks most identifying details +but is not full redaction, so read the bundle before posting it. Two levels are available: + +| Level | How to select | What it masks | +| --- | --- | --- | +| `default` | `-A` / `--anonymize` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept | +| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure | + +See [collecting a debug bundle](https://docs.netbird.io/help/troubleshooting-client#debug-bundle) +and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for) for details. + +See [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075) +for the full workflow, or [SUPPORT.md](SUPPORT.md) for a shorter version. + +### Contributing + +Contributions are welcome. Read [CONTRIBUTING.md](CONTRIBUTING.md) first. NetBird works ticket +first, anything that changes behavior needs an issue the team has agreed on before you open a pull +request. + ### Community projects - [NetBird installer script](https://github.com/physk/netbird-installer) - [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings diff --git a/SUPPORT.md b/SUPPORT.md new file mode 100644 index 000000000..fade37286 --- /dev/null +++ b/SUPPORT.md @@ -0,0 +1,121 @@ +# Getting help with NetBird + +Where to go depends on what you need. If you are not sure, start with +[Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) +and we will move it. + +## Before you post + +1. Search existing [discussions](https://github.com/netbirdio/netbird/discussions) and + [issues](https://github.com/netbirdio/netbird/issues), including closed ones. +2. Check the [documentation](https://docs.netbird.io) and the troubleshooting guides for + [clients](https://docs.netbird.io/help/troubleshooting-client) and + [self-hosted deployments](https://docs.netbird.io/selfhosted/troubleshooting). +3. Remove or anonymize sensitive information from logs, screenshots, and configuration. + +If a discussion already covers your problem, upvote it and add your details there rather than +opening a duplicate. Extra reproduction detail, affected versions, and deployment notes are +useful even on an existing thread. + +## Community support + +Free, for everyone. Covers the NetBird client, open source self-hosted deployments, and general +questions. + +| What you want to do | Where to go | +| --- | --- | +| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) | +| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) | +| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) | +| Chat with the community | [Slack](https://docs.netbird.io/slack-url) | + +## Paid support + +For NetBird Cloud customers and commercial-license self-hosted deployments, covering the +dashboard, control plane, billing, and subscriptions, see +[reporting bugs and issues](https://docs.netbird.io/help/report-bug-issues). + +## Security + +Do not report security vulnerabilities in public issues or discussions, and do not post secrets, +private keys, internal hostnames, or sensitive logs. Use the +[security policy](https://github.com/netbirdio/netbird/security/policy). + +## What makes a report we can act on + +For a bug, the most useful reports include: + +- NetBird version, and component versions where applicable +- Operating system or environment +- Deployment type: NetBird Cloud, self-hosted, Kubernetes, Docker, or local development +- Current behavior and expected behavior +- The smallest set of steps that reproduces the problem +- Logs, status output, screenshots, or a debug bundle when relevant +- Whether this worked before, and the last known working version + +For client reports, these commands usually give us what we need: + +```shell +netbird version +netbird status -d -A +netbird debug for 1m -A -S -U +``` + +`-A` (`--anonymize`) replaces sensitive values consistently across every file in the bundle, so +it stays readable while masking most identifying details. It is not a guarantee of full redaction: +internal address ranges survive at the default level, and interface names, indexes, MTUs, and +flags are never anonymized. Read the bundle before posting it publicly. Two levels are +available: + +| Level | How to select | What it masks | +| --- | --- | --- | +| `default` | `-A` / `--anonymize`, or `--anonymize-level default` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept, and interface names are not anonymized | +| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure | + +Use `strict` when internal addressing or peer naming is itself sensitive. Either way, private +keys and SSH keys are never included, and the packet capture (`capture.pcap`) is left out of +anonymized bundles because it holds raw decrypted packets. + +`-U` (`--upload-bundle`) uploads the bundle and returns a file key you can paste into the thread +instead of attaching an archive. Retention is controlled by the upload service; check its policy +before uploading, and configure cleanup for self-hosted deployments. + +For more detail, see [troubleshooting client issues](https://docs.netbird.io/help/troubleshooting-client), +which explains [what a debug bundle contains](https://docs.netbird.io/help/troubleshooting-client#debug-bundle), +and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for). + +Intermittent problems are still worth reporting. They just need enough detail to investigate: +trigger, frequency, timing, timestamps, and any related logs. + +For a feature request, describe the problem before the solution: what you are trying to +accomplish, who is affected and how often, why the current behavior or workaround is not enough, +and what you would like to see instead. + +## What happens after you post + +Our team, maintainers, or community members may ask for missing details, link related +threads, merge duplicates, move your post to a better category, or try to reproduce the problem. + +Not every discussion becomes an issue. Some are answered in Q&A, some turn out to be +configuration problems, and some need more information before engineering can act. A +well-answered discussion is still a useful outcome. + +When a report is confirmed and actionable, a maintainer opens a validated issue linked back to +the discussion, in whichever repository the fix belongs to. You do not need to know which +repository that is. Routing is part of triage. + +## A note on issues + +Issues in this repository are maintainer-curated work items. Every open issue is something a +maintainer or contributor can pick up and act on. Issues opened without a linked validated +discussion may be closed and redirected here. + +Maintainers can still open issues directly for work found internally, such as regressions caught +during development, planned maintenance, or release blockers. + +## Related reading + +- [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075) +- [Moving to a discussion-first approach](https://github.com/netbirdio/netbird/discussions/6074) +- [CONTRIBUTING.md](CONTRIBUTING.md) for opening pull requests +- [CODE_OF_CONDUCT.md](CODE_OF_CONDUCT.md) diff --git a/agent-network/README.md b/agent-network/README.md index 029ada299..35b9c6668 100644 --- a/agent-network/README.md +++ b/agent-network/README.md @@ -110,12 +110,12 @@ Two roles delegate Agent Network access without account-admin rights: read-only users, groups, peers, and account info (needed to build policies). Nothing else in the account. - **`usage_viewer`** — the regular User baseline plus read on - `agent_network.usage` (the aggregated usage and cost overview) and read-only - access to the resources the usage filters resolve against: users, groups, - peers, and the provider list (connection config redacted — no upstream URLs - or operator-supplied header values). No policies, and no account-wide - request-level access logs; like any caller, it still reads its own requests - through the self-scoped endpoints below. + `agent_network.usage` (the aggregated usage and cost overview) and + `agent_network.logs` (the account-wide request-level access logs, which can + contain captured prompts), and read-only access to the resources those + filters resolve against: users, groups, peers, and the provider list + (connection config redacted — no upstream URLs or operator-supplied header + values). No policies, guardrails, budgets, or settings. Every authenticated user, regardless of role, can read the caller-scoped self-service endpoint `GET /api/agent-network/agent-config` (the endpoint, providers, diff --git a/base62/base62.go b/base62/base62.go index efafbc768..1a02e98e2 100644 --- a/base62/base62.go +++ b/base62/base62.go @@ -3,56 +3,75 @@ package base62 import ( "fmt" "math" - "strings" ) const ( - alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" - base = uint32(len(alphabet)) + alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + base = uint32(len(alphabet)) + maxBase62Digits = 6 // max number of digits required to encode MaxUint32 + ) +var ( + ErrEmptyString = fmt.Errorf("empty string") + ErrInvalidChar = fmt.Errorf("invalid character") + ErrOverflow = fmt.Errorf("integer overflow") +) + +// Fixed-size arrays have better performance and lower memory overhead compared to maps for small static sets of data +var charToIndex [123]int8 // Assuming ASCII from '\0' - 'z' + +func init() { + for i := range charToIndex { + charToIndex[i] = -1 + } + for i, c := range alphabet { + charToIndex[c] = int8(i) + } +} + // Encode encodes a uint32 value to a base62 string. -func Encode(num uint32) string { - if num == 0 { - return string(alphabet[0]) +// The returned string will be between 1-6 characters long. +func Encode(n uint32) string { + if n < base { + return string(alphabet[n]) + } + // avoid dynamic memory usage for small, fixed size data + buf := [maxBase62Digits]byte{} + idx := len(buf) + + for n > 0 { + idx-- + buf[idx] = alphabet[n%base] + n /= base } - var encoded strings.Builder - - for num > 0 { - remainder := num % base - encoded.WriteByte(alphabet[remainder]) - num /= base - } - - // Reverse the encoded string - encodedString := encoded.String() - reversed := reverse(encodedString) - return reversed + return string(buf[idx:]) } // Decode decodes a base62 string to a uint32 value. +// Returns an error if the input string is empty, contains invalid characters, +// or would result in integer overflow. func Decode(encoded string) (uint32, error) { + if len(encoded) == 0 { + return 0, ErrEmptyString + } var decoded uint32 - strLen := len(encoded) - - for i, char := range encoded { - index := strings.IndexRune(alphabet, char) + for _, char := range encoded { + index := int8(-1) + if int(char) < len(charToIndex) { + index = charToIndex[char] + } if index < 0 { - return 0, fmt.Errorf("invalid character: %c", char) + return 0, fmt.Errorf("%w: %c", ErrInvalidChar, char) + } + // Add overflow check when calculating the decoded value to prevent silent overflow of uint32 + if decoded > (math.MaxUint32-uint32(index))/base { + return 0, fmt.Errorf("%w: %s", ErrOverflow, encoded) } - decoded += uint32(index) * uint32(math.Pow(float64(base), float64(strLen-i-1))) + decoded = decoded*base + uint32(index) } return decoded, nil } - -// Reverse a string. -func reverse(s string) string { - runes := []rune(s) - for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 { - runes[i], runes[j] = runes[j], runes[i] - } - return string(runes) -} diff --git a/base62/base62_test.go b/base62/base62_test.go index 00da2124a..f2ad06d6f 100644 --- a/base62/base62_test.go +++ b/base62/base62_test.go @@ -1,31 +1,67 @@ package base62 import ( + "errors" + "math" "testing" ) func TestEncodeDecode(t *testing.T) { - tests := []struct { - num uint32 + testCases := []struct { + input uint32 + expected string }{ - {0}, - {1}, - {42}, - {12345}, - {99999}, - {123456789}, + {0, "0"}, + {1, "1"}, + {5, "5"}, + {9, "9"}, + {10, "A"}, + {42, "g"}, + {61, "z"}, + {62, "10"}, + {'0', "m"}, + {'9', "v"}, + {'A', "13"}, + {'Z', "1S"}, + {'a', "1Z"}, + {'z', "1y"}, + {99999, "Q0t"}, + {12345, "3D7"}, + {123456789, "8M0kX"}, + {math.MaxUint32, "4gfFC3"}, } - for _, tt := range tests { - encoded := Encode(tt.num) + for _, tc := range testCases { + encoded := Encode(tc.input) + if encoded != tc.expected { + t.Errorf("Encode(%d) = %s; want %s", tc.input, encoded, tc.expected) + } decoded, err := Decode(encoded) - if err != nil { - t.Errorf("Decode error: %v", err) + t.Errorf("Expected error nil, got %v", err) } - if decoded != tt.num { - t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tt.num) + if decoded != tc.input { + t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tc.input) } } } + +// Decode handles empty string input with appropriate error +func TestDecodeEmptyString(t *testing.T) { + if _, err := Decode(""); !errors.Is(err, ErrEmptyString) { + t.Errorf("Expected error %v, got %v", ErrEmptyString, err) + } +} + +func TestDecodeOverflow(t *testing.T) { + if _, err := Decode("4gfFC4"); !errors.Is(err, ErrOverflow) { + t.Errorf("Expected error %v, got %v", ErrOverflow, err) + } +} + +func TestDecodeInvalid(t *testing.T) { + if _, err := Decode("/"); !errors.Is(err, ErrInvalidChar) { + t.Errorf("Expected error %v, got %v", ErrInvalidChar, err) + } +} diff --git a/client/android/client.go b/client/android/client.go index e47a1c13d..9705db8e0 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -104,8 +104,7 @@ type Client struct { stateChangeMu sync.Mutex stateChangeSubID string - eventSub *peer.EventSubscription - // Closed to stop the watch goroutines from delivering buffered items to a + // Closed to stop the watch goroutine from delivering buffered ticks to a // listener that has been removed or replaced. See stopStateChangeWatchLocked. stateChangeDone chan struct{} @@ -213,6 +212,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetEvents(c.netMgr)) c.setState(cfg, cacheDir, cfgFile, connectClient) + connectClient.SetSyncResponsePersistence(true) // This path runs the interactive SSO flow, so reaching here means the peer // is authenticated again — release the latch Status() reports from. Clear // only once the fresh connect client is installed: until then Status() @@ -256,6 +256,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetEvents(c.netMgr)) c.setState(cfg, cacheDir, cfgFile, connectClient) + connectClient.SetSyncResponsePersistence(true) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } @@ -327,6 +328,19 @@ func (c *Client) NotifyNetworkChange() { // or "strict"; strict also anonymizes internal IP ranges, peer names, and // WireGuard public keys, and implies anonymize. func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) { + return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true) +} + +// DebugBundleFile generates a debug bundle and returns the path of the zip in +// the cache directory instead of uploading it, so the app can hand the file to +// the user for inspection. The caller owns the file and removes it once done; +// the stale-bundle cleanup of later runs removes it only after a day. +// anonymize and anonymizeLevel behave as in DebugBundle. +func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) { + return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false) +} + +func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) { cfg, cacheDir, cc := c.stateSnapshot() // If the engine hasn't been started, load config from disk @@ -342,6 +356,11 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym cacheDir = platformFiles.CacheDir() } + // Clear what an interrupted earlier run may have left in the cache before + // adding to it. Remote debug jobs write to the same directory, so anything + // younger than an hour is treated as possibly still in use. + debug.RemoveStaleBundles(cacheDir, time.Hour) + deps := debug.GeneratorDependencies{ InternalConfig: cfg, StatusRecorder: c.recorder, @@ -379,6 +398,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym if err != nil { return "", fmt.Errorf("generate debug bundle: %w", err) } + if !upload { + return debug.ExportBundle(path) + } defer func() { if err := os.Remove(path); err != nil { log.Errorf("failed to remove debug bundle file: %v", err) @@ -475,6 +497,7 @@ func (c *Client) Networks() *NetworkArray { routesMap := routeManager.GetClientRoutesWithNetID() v6Merged := route.V6ExitMergeSet(routesMap) resolvedDomains := c.recorder.GetResolvedDomainsStates() + activeRoutePeers := c.recorder.GetActiveRoutePeers() networkArray := &NetworkArray{ items: make([]Network, 0), @@ -488,7 +511,7 @@ func (c *Client) Networks() *NetworkArray { continue } - network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged) + network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers) if network == nil { continue } @@ -497,14 +520,14 @@ func (c *Client) Networks() *NetworkArray { return networkArray } -func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network { +func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network { r := routes[0] netStr := r.Network.String() if r.IsDynamic() { netStr = r.Domains.SafeString() } - routePeer, err := c.findBestRoutePeer(routes) + routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers) if err != nil { log.Errorf("could not get peer info for route %s: %v", id, err) return nil @@ -528,12 +551,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo // findBestRoutePeer returns the peer actively routing traffic for the given // HA route group. Falls back to the first connected peer, then the first peer. -func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) { - netStr := routes[0].Network.String() - - fullStatus := c.recorder.GetFullStatus() - for _, p := range fullStatus.Peers { - if _, ok := p.GetRoutes()[netStr]; ok { +func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) { + if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok { + if p, err := c.recorder.GetPeer(peerKey); err == nil { return p, nil } } diff --git a/client/android/session.go b/client/android/session.go index d5da09c93..1ce97f074 100644 --- a/client/android/session.go +++ b/client/android/session.go @@ -6,13 +6,8 @@ import ( "context" "fmt" - log "github.com/sirupsen/logrus" - "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/auth" - "github.com/netbirdio/netbird/client/internal/auth/sessionwatch" - "github.com/netbirdio/netbird/client/internal/peer" - cProto "github.com/netbirdio/netbird/client/proto" ) // StateChangeListener receives client state notifications. @@ -21,16 +16,11 @@ import ( // changed: connection state, the run-loop status label (e.g. NeedsLogin) or // the session deadline. It mirrors the daemon's SubscribeStatus stream // trigger — on each signal the consumer pulls the fresh values via -// Status() / SessionExpiresAtUnix(). -// -// OnSessionExpiring forwards the engine's session-expiry warnings, fired at -// sessionwatch.WarningLead before the deadline and again at FinalWarningLead -// (finalWarning true). The second one is suppressed when the user dismissed -// the first via DismissSessionWarning. The daemon turns the same events into -// its tray notification. +// Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning +// timers on Android; the app schedules the warnings from the deadline it +// reads here. type StateChangeListener interface { OnStateChanged() - OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool) } // Status returns the connect run-loop's status label — the same value the @@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) { return } - // Both subscriptions are buffered (one pending tick, ten pending events), - // so unsubscribing is not enough to stop callbacks: the loops would drain - // what is already queued and deliver it to a listener the caller has - // already removed or replaced. Gate every callback on this registration's - // own signal, which is closed before unsubscribing. + // The subscription is buffered (one pending tick), so unsubscribing is + // not enough to stop callbacks: the loop would drain what is already + // queued and deliver it to a listener the caller has already removed or + // replaced. Gate every callback on this registration's own signal, which + // is closed before unsubscribing. done := make(chan struct{}) c.stateChangeDone = done @@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) { listener.OnStateChanged() } }() - - c.eventSub = c.recorder.SubscribeToEvents() - go watchSessionWarnings(c.eventSub, listener, done) } // RemoveStateChangeListener unregisters the state notification listener. @@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() { c.stopStateChangeWatchLocked() } -// DismissSessionWarning records the user's "Dismiss" on the first expiry -// warning and suppresses the final one for the current deadline. A refreshed -// deadline re-arms both. No-op while the engine is not running. -func (c *Client) DismissSessionWarning() { - cc := c.getConnectClient() - if cc == nil { - return - } - engine := cc.Engine() - if engine == nil { - return - } - engine.DismissSessionWarning() -} - // ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and // asks the management server to extend the session deadline. The tunnel is // untouched: no resync, no reconnect. Async; the result arrives on the @@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() { } func (c *Client) stopStateChangeWatchLocked() { - // Signal first, unsubscribe second: closing the channels only stops new - // items, and the loops would still hand whatever is buffered to a listener + // Signal first, unsubscribe second: closing the channel only stops new + // items, and the loop would still hand whatever is buffered to a listener // that is no longer registered. if c.stateChangeDone != nil { close(c.stateChangeDone) @@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() { c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID) c.stateChangeSubID = "" } - if c.eventSub != nil { - // Closes the channel, which ends watchSessionWarnings. - c.recorder.UnsubscribeFromEvents(c.eventSub) - c.eventSub = nil - } -} - -// watchSessionWarnings forwards the engine's session-expiry warnings to the -// listener. The event stream also carries unrelated traffic — network-map -// updates on every sync, DNS and route errors — so everything but an -// AUTHENTICATION event carrying the session-warning marker is dropped. Exits -// when the subscription is closed by UnsubscribeFromEvents, or earlier when -// done is closed — the stream buffers up to ten events, and a deregistered -// listener must not receive the ones already queued. -func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) { - for ev := range sub.Events() { - select { - case <-done: - return - default: - } - if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION { - continue - } - meta := ev.GetMetadata() - if meta[sessionwatch.MetaSessionWarning] != "true" { - // Other AUTHENTICATION events exist (e.g. a deadline rejected as - // out of range); they carry no warning marker. - continue - } - deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt]) - if err != nil { - log.Warnf("session warning event with unparsable deadline: %v", err) - continue - } - lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes]) - if err != nil { - // Informational only — the deadline above is what drives the UI. - lead = 0 - } - listener.OnSessionExpiring(deadline.Unix(), int64(lead), - meta[sessionwatch.MetaSessionFinal] == "true") - } } func (c *Client) beginExtend() (context.Context, error) { diff --git a/client/android/ssh_client.go b/client/android/ssh_client.go index 2822b6539..9a11044ba 100644 --- a/client/android/ssh_client.go +++ b/client/android/ssh_client.go @@ -31,6 +31,8 @@ const ( // PasswordRequiredMarker tells Java to prompt for a password and retry. It is // a string because gomobile flattens errors to their message, so a sentinel // value would not survive the binding. +// +//nolint:gosec // G101 false positive: a sentinel marker, not a credential const PasswordRequiredMarker = "netbird-ssh-password-required" // HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation, diff --git a/client/cmd/debug.go b/client/cmd/debug.go index 98fe53626..c4d5ad6d7 100644 --- a/client/cmd/debug.go +++ b/client/cmd/debug.go @@ -23,7 +23,10 @@ import ( "github.com/netbirdio/netbird/version" ) -const errCloseConnection = "Failed to close connection: %v" +const ( + errCloseConnection = "Failed to close connection: %v" + noUpDownFlag = "no-updown" +) var ( logFileCount uint32 @@ -257,13 +260,14 @@ func runForDuration(cmd *cobra.Command, args []string) error { } stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting) + noUpDown, _ := cmd.Flags().GetBool(noUpDownFlag) initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{}) if err != nil { return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message()) } - if stateWasDown { + if stateWasDown && !noUpDown { if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil { cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message()) } else { @@ -284,34 +288,20 @@ func runForDuration(cmd *cobra.Command, args []string) error { } needsRestoreUp := false - if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil { - cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message()) + if noUpDown { + enableSyncResponsePersistence(cmd, client) } else { - needsRestoreUp = !stateWasDown - cmd.Println("netbird down") + needsRestoreUp = restartDaemon(cmd, client, stateWasDown) } - time.Sleep(1 * time.Second) - - // Enable sync response persistence before bringing the service up - if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{ - Enabled: true, - }); err != nil { - cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message()) - } - - if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil { - cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message()) - } else { - needsRestoreUp = false - cmd.Println("netbird up") - } - - time.Sleep(3 * time.Second) - cpuProfilingStarted := false if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil { - cmd.PrintErrf("Failed to start CPU profiling: %v\n", err) + if msg := status.Convert(err).Message(); strings.Contains(msg, "already in progress") { + cmd.PrintErrln("CPU profiling is already running (started with `netbird debug cpu start`). " + + "It is left running and is included in a bundle created after `netbird debug cpu stop`.") + } else { + cmd.PrintErrf("Failed to start CPU profiling: %v\n", msg) + } } else { cpuProfilingStarted = true defer func() { @@ -401,7 +391,7 @@ func runForDuration(cmd *cobra.Command, args []string) error { } } - if stateWasDown { + if stateWasDown && !noUpDown { if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil { cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message()) } else { @@ -458,6 +448,47 @@ func setSyncResponsePersistence(cmd *cobra.Command, args []string) error { return nil } +// enableSyncResponsePersistence asks the daemon to keep the latest sync +// response so the bundle carries the network map. With a running daemon only +// syncs received after the call are kept. +func enableSyncResponsePersistence(cmd *cobra.Command, client proto.DaemonServiceClient) { + if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{ + Enabled: true, + }); err != nil { + cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message()) + } +} + +// restartDaemon cycles the daemon down and up with sync response persistence +// enabled so the bundle carries the network map. It reports whether the +// daemon was left down although it was running before, so the caller can +// bring it back up. +func restartDaemon(cmd *cobra.Command, client proto.DaemonServiceClient, stateWasDown bool) bool { + needsRestoreUp := false + if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil { + cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message()) + } else { + needsRestoreUp = !stateWasDown + cmd.Println("netbird down") + } + + time.Sleep(1 * time.Second) + + // Enable sync response persistence before bringing the service up + enableSyncResponsePersistence(cmd, client) + + if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil { + cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message()) + } else { + needsRestoreUp = false + cmd.Println("netbird up") + } + + time.Sleep(3 * time.Second) + + return needsRestoreUp +} + func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error { ticker := time.NewTicker(1 * time.Second) defer ticker.Stop() @@ -546,4 +577,5 @@ func init() { forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle") forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root") forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle") + forCmd.Flags().Bool(noUpDownFlag, false, "Keep the daemon running instead of bringing it down and up before collecting. The bundle only includes the network map if a sync arrives during the run") } diff --git a/client/cmd/debug_cpu.go b/client/cmd/debug_cpu.go new file mode 100644 index 000000000..a01b845cf --- /dev/null +++ b/client/cmd/debug_cpu.go @@ -0,0 +1,83 @@ +package cmd + +import ( + "fmt" + + log "github.com/sirupsen/logrus" + "github.com/spf13/cobra" + "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/proto" +) + +var debugCPUCmd = &cobra.Command{ + Use: "cpu", + Short: "Profile the daemon's CPU usage", + Long: `Starts and stops CPU profiling in the running daemon without restarting it. +The profile is included in the next debug bundle as cpu.prof. + +Profiling is not time limited: it keeps running, and keeps costing CPU, until +"netbird debug cpu stop" is run.`, +} + +var debugCPUStartCmd = &cobra.Command{ + Use: "start", + Short: "Start CPU profiling in the daemon", + Example: " netbird debug cpu start", + Args: cobra.NoArgs, + RunE: debugCPUStart, +} + +var debugCPUStopCmd = &cobra.Command{ + Use: "stop", + Short: "Stop CPU profiling in the daemon", + Long: `Stops CPU profiling. The captured profile stays in the daemon until the next +debug bundle is created, which includes it as cpu.prof.`, + Example: " netbird debug cpu stop && netbird debug bundle", + Args: cobra.NoArgs, + RunE: debugCPUStop, +} + +func debugCPUStart(cmd *cobra.Command, _ []string) error { + conn, err := getClient(cmd) + if err != nil { + return err + } + defer func() { + if err := conn.Close(); err != nil { + log.Errorf(errCloseConnection, err) + } + }() + + if _, err := proto.NewDaemonServiceClient(conn).StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil { + return fmt.Errorf("start CPU profiling: %v", status.Convert(err).Message()) + } + + cmd.Println("CPU profiling started and runs until stopped. Run `netbird debug cpu stop` and then `netbird debug bundle` to collect it.") + return nil +} + +func debugCPUStop(cmd *cobra.Command, _ []string) error { + conn, err := getClient(cmd) + if err != nil { + return err + } + defer func() { + if err := conn.Close(); err != nil { + log.Errorf(errCloseConnection, err) + } + }() + + if _, err := proto.NewDaemonServiceClient(conn).StopCPUProfile(cmd.Context(), &proto.StopCPUProfileRequest{}); err != nil { + return fmt.Errorf("stop CPU profiling: %v", status.Convert(err).Message()) + } + + cmd.Println("CPU profiling stopped. Run `netbird debug bundle` to include cpu.prof.") + return nil +} + +func init() { + debugCPUCmd.AddCommand(debugCPUStartCmd) + debugCPUCmd.AddCommand(debugCPUStopCmd) + debugCmd.AddCommand(debugCPUCmd) +} diff --git a/client/cmd/debug_cpu_test.go b/client/cmd/debug_cpu_test.go new file mode 100644 index 000000000..85fffd462 --- /dev/null +++ b/client/cmd/debug_cpu_test.go @@ -0,0 +1,164 @@ +package cmd + +import ( + "bytes" + "context" + "os/user" + "strings" + "testing" + + "github.com/spf13/cobra" + "github.com/spf13/pflag" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/internal/profilemanager" +) + +// startDebugTestDaemon starts an in-process daemon with an isolated profile +// directory and returns the address the CLI should dial. +func startDebugTestDaemon(t *testing.T) string { + t.Helper() + + tempDir := t.TempDir() + origDefaultProfileDir := profilemanager.DefaultConfigPathDir + origActiveProfileStatePath := profilemanager.ActiveProfileStatePath + origConfigDirOverride := profilemanager.ConfigDirOverride + origDaemonAddr := daemonAddr + t.Cleanup(func() { + profilemanager.DefaultConfigPathDir = origDefaultProfileDir + profilemanager.ActiveProfileStatePath = origActiveProfileStatePath + profilemanager.ConfigDirOverride = origConfigDirOverride + daemonAddr = origDaemonAddr + }) + + profilemanager.DefaultConfigPathDir = tempDir + profilemanager.ActiveProfileStatePath = tempDir + "/active_profile.json" + profilemanager.ConfigDirOverride = tempDir + + currUser, err := user.Current() + require.NoError(t, err) + sm := profilemanager.ServiceManager{} + created, err := sm.AddProfile("test1", currUser.Username) + require.NoError(t, err) + require.NoError(t, sm.SetActiveProfileState(&profilemanager.ActiveProfileState{ + ID: created.ID, + Username: currUser.Username, + })) + + ctx, cancel := context.WithCancel(internal.CtxInitState(context.Background())) + srv, lis := startClientDaemon(t, ctx, "", tempDir+"/config.json") + t.Cleanup(func() { + cancel() + srv.Stop() + }) + + return "tcp://" + lis.Addr().String() +} + +// runDebugCmd runs `netbird debug ` against the daemon at addr and +// returns everything the command printed. +func runDebugCmd(addr string, args ...string) (string, error) { + daemonAddr = addr + var out bytes.Buffer + rootCmd.SetOut(&out) + rootCmd.SetErr(&out) + rootCmd.SetArgs(append(append([]string{"debug"}, args...), "--daemon-addr", addr, "--log-file", "")) + err := rootCmd.Execute() + rootCmd.SetOut(nil) + rootCmd.SetErr(nil) + rootCmd.SetArgs(nil) + resetFlags(rootCmd) + return out.String(), err +} + +// resetFlags puts every flag of the command and its subcommands back to its +// default so a value parsed in one run does not leak into the next in-process +// execution. +func resetFlags(cmd *cobra.Command) { + reset := func(f *pflag.Flag) { + // Set appends to a slice flag and would parse the "[a,b]" default + // text as elements, so slices are replaced instead. + if sv, ok := f.Value.(pflag.SliceValue); ok { + var def []string + if trimmed := strings.Trim(f.DefValue, "[]"); trimmed != "" { + def = strings.Split(trimmed, ",") + } + _ = sv.Replace(def) + } else { + _ = f.Value.Set(f.DefValue) + } + f.Changed = false + } + cmd.Flags().VisitAll(reset) + cmd.PersistentFlags().VisitAll(reset) + // Commands pin their writers to the buffer of the run that first used + // them, so a later run would print into the old buffer. + cmd.SetOut(nil) + cmd.SetErr(nil) + for _, sub := range cmd.Commands() { + resetFlags(sub) + } +} + +// TestResetFlagsSliceDefault guards against Set("[]") on slice flags, which +// stores a literal "[]" element instead of the empty default. +func TestResetFlagsSliceDefault(t *testing.T) { + cmd := &cobra.Command{Use: "x"} + var env, withDefault []string + cmd.Flags().StringSliceVar(&env, "env", nil, "") + cmd.Flags().StringSliceVar(&withDefault, "names", []string{"a", "b"}, "") + require.NoError(t, cmd.Flags().Parse([]string{"--env", "K=V", "--names", "c"})) + + resetFlags(cmd) + + assert.Empty(t, env, "slice flag with no default must reset to empty") + assert.Equal(t, []string{"a", "b"}, withDefault, "slice flag must reset to its default") +} + +func TestDebugCPUStartStop(t *testing.T) { + addr := startDebugTestDaemon(t) + + run := func(args ...string) error { + _, err := runDebugCmd(addr, append([]string{"cpu"}, args...)...) + return err + } + + require.Error(t, run("stop"), "stop without a running profile must fail") + require.NoError(t, run("start")) + assert.Error(t, run("start"), "second start must be rejected while profiling") + require.NoError(t, run("stop")) + assert.Error(t, run("stop"), "second stop must be rejected") + assert.NoError(t, run("start"), "profiling can be started again after a stop") + assert.NoError(t, run("stop")) +} + +// TestDebugForKeepsRunningCPUProfile covers `debug for` started while a +// profile from `debug cpu start` is running: it must say so, leave the +// profile alone, and still create the bundle. +func TestDebugForKeepsRunningCPUProfile(t *testing.T) { + addr := startDebugTestDaemon(t) + + _, err := runDebugCmd(addr, "cpu", "start") + require.NoError(t, err) + + out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown") + require.NoError(t, err, "output: %s", out) + assert.Contains(t, out, "CPU profiling is already running", "the conflict must be explained") + assert.NotContains(t, out, "rpc error", "the raw RPC error must not reach the user") + assert.Contains(t, out, "Local file:", "the bundle must still be created") + + _, err = runDebugCmd(addr, "cpu", "stop") + assert.NoError(t, err, "the profile started by the user must still be running") +} + +func TestDebugForNoUpDown(t *testing.T) { + addr := startDebugTestDaemon(t) + + out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown") + require.NoError(t, err, "output: %s", out) + assert.NotContains(t, out, "netbird down", "--no-updown must not bring the daemon down") + assert.NotContains(t, out, "netbird up", "--no-updown must not bring the daemon up") + assert.Contains(t, out, "Local file:", "the bundle must still be created") +} diff --git a/client/cmd/service.go b/client/cmd/service.go index 7410d60ea..2a558e6d5 100644 --- a/client/cmd/service.go +++ b/client/cmd/service.go @@ -7,6 +7,7 @@ import ( "fmt" "net/http" "runtime" + "slices" "strings" "sync" @@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{ const defaultJSONSocket = "unix:///var/run/netbird-http.sock" +// forbiddenServiceEnvVars are the environment variables the service is never +// registered with, keyed in upper case since these are Windows names. Each one +// decides where the daemon resolves something it then uses with the privileges +// of the account it runs under — LocalSystem on Windows, root elsewhere: the +// executables it runs (PATH, PATHEXT, COMSPEC, SystemRoot, windir) or the +// directory it writes temporary files in (TEMP, TMP). The daemon needs none of +// them, and the utilities it shells out to are resolved by absolute path. +var forbiddenServiceEnvVars = map[string]struct{}{ + "PATH": {}, + "PATHEXT": {}, + "SYSTEMROOT": {}, + "WINDIR": {}, + "COMSPEC": {}, + "TEMP": {}, + "TMP": {}, +} + +// forbiddenServiceEnvPrefixes are the dynamic-loader families, refused whole +// rather than by name: LD_PRELOAD, DYLD_INSERT_LIBRARIES and their siblings all +// reach the loader of the process, the set differs per platform and libc, and +// new members arrive with new OS releases. Listing them one by one is a list +// that is wrong the moment it is written. +var forbiddenServiceEnvPrefixes = []string{"LD_", "DYLD_"} + var ( serviceName string serviceEnvVars []string @@ -127,8 +152,33 @@ func parseServiceEnvVars(envVars []string) (map[string]string, error) { return nil, fmt.Errorf("empty environment variable key in: %s", env) } + if isForbiddenServiceEnvVar(key) { + return nil, fmt.Errorf("environment variable %s cannot be set on the service: it decides where the service resolves the executables, libraries or temporary files it uses", key) + } + envMap[key] = value } return envMap, nil } + +// isForbiddenServiceEnvVar reports whether name is one the service must not be +// registered with. +// +// The names are matched case-insensitively only on Windows, where they are the +// same variable however they are spelled. Elsewhere the environment is +// case-sensitive, so Path and PATH are two different variables and only the +// exact spelling is the one the loader reads. +func isForbiddenServiceEnvVar(name string) bool { + if runtime.GOOS == "windows" { + name = strings.ToUpper(name) + } + + if _, forbidden := forbiddenServiceEnvVars[name]; forbidden { + return true + } + + return slices.ContainsFunc(forbiddenServiceEnvPrefixes, func(prefix string) bool { + return strings.HasPrefix(name, prefix) + }) +} diff --git a/client/cmd/service_params.go b/client/cmd/service_params.go index 750b22ae6..6e2dbec40 100644 --- a/client/cmd/service_params.go +++ b/client/cmd/service_params.go @@ -14,6 +14,7 @@ import ( "github.com/netbirdio/netbird/client/configs" "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/elevate" "github.com/netbirdio/netbird/util" ) @@ -43,10 +44,33 @@ func serviceParamsPath() string { // loadServiceParams reads saved service parameters from disk. // Returns nil with no error if the file does not exist. +// +// The file is read by an elevated install and decides the arguments and the +// environment of the service it then registers, so it is used only when its +// ownership and permissions are the ones saveServiceParams leaves behind. That +// restricted ACL is applied when the file is written, which is not necessarily +// before it is first read, so this is checked rather than assumed. A file that +// fails the check is treated as absent, and the install proceeds with its +// defaults. func loadServiceParams() (*serviceParams, error) { path := serviceParamsPath() - data, err := os.ReadFile(path) + // Resolve links first so the checks apply to the file that is actually read. + // Since the check covers every directory above it as well, nobody who fails + // it can swap the file between here and the read below. + resolved, err := filepath.EvalSymlinks(path) + if err != nil { + if os.IsNotExist(err) { + return nil, nil //nolint:nilnil + } + return nil, fmt.Errorf("resolve service params %s: %w", path, err) + } + + if err := elevate.CheckOnlyOwnerWritable(resolved); err != nil { + return nil, fmt.Errorf("refusing to read service params from %s: %w", resolved, err) + } + + data, err := os.ReadFile(resolved) if err != nil { if os.IsNotExist(err) { return nil, nil //nolint:nilnil @@ -182,10 +206,16 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) { // If --service-env was explicitly set to empty, all saved env vars are cleared. // If --service-env was not set, saved env vars are used entirely. func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) { + // A forbidden name explicitly passed on the command line is an error the + // operator is told about, but one restored from a file written by an older + // version is dropped: an install that refuses to run would leave the host + // without a daemon over a variable nobody is asking for any more. + saved := dropForbiddenServiceEnvVars(cmd, params.ServiceEnvVars) + if !cmd.Flags().Changed("service-env") { - if len(params.ServiceEnvVars) > 0 { + if len(saved) > 0 { // No explicit env vars: rebuild serviceEnvVars from saved params. - serviceEnvVars = envMapToSlice(params.ServiceEnvVars) + serviceEnvVars = envMapToSlice(saved) } return } @@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) { return } - if len(params.ServiceEnvVars) == 0 { + if len(saved) == 0 { return } // Merge saved values underneath explicit ones. - merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit)) - maps.Copy(merged, params.ServiceEnvVars) + merged := make(map[string]string, len(saved)+len(explicit)) + maps.Copy(merged, saved) maps.Copy(merged, explicit) // explicit wins on conflict serviceEnvVars = envMapToSlice(merged) } @@ -233,6 +263,20 @@ var resetParamsCmd = &cobra.Command{ }, } +// dropForbiddenServiceEnvVars returns the saved entries that may still be +// registered on the service, reporting every one it leaves behind. +func dropForbiddenServiceEnvVars(cmd *cobra.Command, saved map[string]string) map[string]string { + kept := make(map[string]string, len(saved)) + for key, value := range saved { + if isForbiddenServiceEnvVar(key) { + cmd.PrintErrf("Warning: ignoring saved service environment variable %s: it decides where the service resolves the executables, libraries or temporary files it uses\n", key) + continue + } + kept[key] = value + } + return kept +} + // envMapToSlice converts a map of env vars to a KEY=VALUE slice. func envMapToSlice(m map[string]string) []string { s := make([]string, 0, len(m)) diff --git a/client/cmd/service_params_test.go b/client/cmd/service_params_test.go index 94f98a0ce..1f83374cb 100644 --- a/client/cmd/service_params_test.go +++ b/client/cmd/service_params_test.go @@ -9,6 +9,7 @@ import ( "go/token" "os" "path/filepath" + "runtime" "strings" "testing" @@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) { assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result) } +func TestParseServiceEnvVars_RejectsForbiddenNames(t *testing.T) { + for _, env := range []string{"PATH=C:\\somewhere", "LD_PRELOAD=/tmp/lib.so", "DYLD_FALLBACK_LIBRARY_PATH=/tmp"} { + _, err := parseServiceEnvVars([]string{"KEEP=me", env}) + require.Errorf(t, err, "%s selects what the service resolves and must be refused", env) + } +} + +func TestIsForbiddenServiceEnvVar(t *testing.T) { + // The loader families are matched by prefix, so a name nobody has heard of + // yet is refused too. + for _, name := range []string{ + "PATH", "PATHEXT", "COMSPEC", "SYSTEMROOT", "WINDIR", "TEMP", "TMP", + "LD_PRELOAD", "LD_AUDIT", "DYLD_INSERT_LIBRARIES", "DYLD_FALLBACK_FRAMEWORK_PATH", + } { + assert.Truef(t, isForbiddenServiceEnvVar(name), "%s must be refused", name) + } + + // The prefix must not swallow names that merely start with the same letters. + for _, name := range []string{"NB_LOG_LEVEL", "NB_WG_DEBUG", "HTTPS_PROXY", "LDAP_URL", "DYLDX"} { + assert.Falsef(t, isForbiddenServiceEnvVar(name), "%s has no reason to be refused", name) + } + + // On Windows a variable is the same one however it is spelled; elsewhere + // Path and PATH are two variables and only the exact one is read. + if runtime.GOOS == "windows" { + assert.True(t, isForbiddenServiceEnvVar("Path")) + assert.True(t, isForbiddenServiceEnvVar("ld_preload")) + } else { + assert.False(t, isForbiddenServiceEnvVar("Path")) + assert.False(t, isForbiddenServiceEnvVar("ld_preload")) + } +} + +func TestApplyServiceEnvParams_DropsForbiddenSavedNames(t *testing.T) { + origServiceEnvVars := serviceEnvVars + t.Cleanup(func() { serviceEnvVars = origServiceEnvVars }) + + serviceEnvVars = nil + + cmd := &cobra.Command{} + cmd.Flags().StringSlice("service-env", nil, "") + + saved := &serviceParams{ + ServiceEnvVars: map[string]string{"PATH": "C:\\attacker", "NB_LOG_FORMAT": "json"}, + } + + applyServiceEnvParams(cmd, saved) + + result, err := parseServiceEnvVars(serviceEnvVars) + require.NoError(t, err, "a saved PATH must be dropped rather than fail the install") + assert.Equal(t, map[string]string{"NB_LOG_FORMAT": "json"}, result) +} + func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) { origServiceEnvVars := serviceEnvVars t.Cleanup(func() { serviceEnvVars = origServiceEnvVars }) diff --git a/client/cmd/service_params_trust_test.go b/client/cmd/service_params_trust_test.go new file mode 100644 index 000000000..1cf564445 --- /dev/null +++ b/client/cmd/service_params_trust_test.go @@ -0,0 +1,57 @@ +//go:build !windows && !ios && !android + +package cmd + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/configs" +) + +// The Windows equivalent of this is the ACL check in +// elevate.CheckOnlyOwnerWritable, covered by that package's own tests; here the +// point is that loadServiceParams asks the question at all. +func TestLoadServiceParams_RefusesWorldWritableFile(t *testing.T) { + tmpDir := t.TempDir() + + original := configs.StateDir + t.Cleanup(func() { configs.StateDir = original }) + configs.StateDir = tmpDir + + path := filepath.Join(tmpDir, serviceParamsFile) + require.NoError(t, os.WriteFile(path, []byte(`{"log_level":"debug"}`), 0o666)) + // WriteFile is subject to the umask, so set the bits that matter explicitly. + require.NoError(t, os.Chmod(path, 0o666)) + + params, err := loadServiceParams() + require.Error(t, err, "a service.json anyone can rewrite must not be trusted") + assert.Nil(t, params) + + require.NoError(t, os.Chmod(path, 0o600)) + params, err = loadServiceParams() + require.NoError(t, err) + require.NotNil(t, params) + assert.Equal(t, "debug", params.LogLevel) +} + +func TestLoadServiceParams_RefusesWorldWritableDirectory(t *testing.T) { + tmpDir := t.TempDir() + stateDir := filepath.Join(tmpDir, "state") + require.NoError(t, os.Mkdir(stateDir, 0o777)) + require.NoError(t, os.Chmod(stateDir, 0o777)) + + original := configs.StateDir + t.Cleanup(func() { configs.StateDir = original }) + configs.StateDir = stateDir + + require.NoError(t, os.WriteFile(filepath.Join(stateDir, serviceParamsFile), []byte(`{}`), 0o600)) + + params, err := loadServiceParams() + require.Error(t, err, "a service.json in a directory anyone can replace entries in must not be trusted") + assert.Nil(t, params) +} diff --git a/client/firewall/iptables/family_linux.go b/client/firewall/iptables/family_linux.go index c5ed8cc20..0e1ce5440 100644 --- a/client/firewall/iptables/family_linux.go +++ b/client/firewall/iptables/family_linux.go @@ -24,6 +24,7 @@ const ( tableFilter = "filter" tableNat = "nat" tableMangle = "mangle" + tableRaw = "raw" // chainACLInput is the peer ACL chain that holds installed // peer-filtering rules. @@ -34,6 +35,7 @@ const ( mangleForwardKey chainKey = "MANGLE-FORWARD" chainInput = "INPUT" + chainOutput = "OUTPUT" chainPostrouting = "POSTROUTING" chainPrerouting = "PREROUTING" chainForward = "FORWARD" diff --git a/client/firewall/iptables/manager_linux.go b/client/firewall/iptables/manager_linux.go index 49b88f1ea..0f0b0110e 100644 --- a/client/firewall/iptables/manager_linux.go +++ b/client/firewall/iptables/manager_linux.go @@ -25,9 +25,8 @@ type Manager struct { wgIface iFaceMapper - ipv4Client *iptables.IPTables - family4 *family - rawSupported bool + ipv4Client *iptables.IPTables + family4 *family // IPv6 counterparts, nil when no v6 overlay ipv6Client *iptables.IPTables @@ -108,10 +107,6 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error { return err } - if err := m.initNoTrackChain(); err != nil { - log.Warnf("raw table not available, notrack rules will be disabled: %v", err) - } - // Trust after all fatal init steps so a later failure doesn't leave the // interface in firewalld's trusted zone without a corresponding Close. if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil { @@ -285,10 +280,6 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error { var merr *multierror.Error - if err := m.cleanupNoTrackChain(); err != nil { - merr = multierror.Append(merr, fmt.Errorf("cleanup notrack chain: %w", err)) - } - if m.hasIPv6() { if err := m.family6.Reset(); err != nil { merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err)) @@ -440,134 +431,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort) } -const ( - chainNameRaw = "NETBIRD-RAW" - chainOutput = "OUTPUT" - tableRaw = "raw" -) - -// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic. -// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which -// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark). -// -// Traffic flows that need NOTRACK: -// -// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite) -// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort -// Matched by: sport=wgPort -// -// 2. Egress: Proxy -> WireGuard (via raw socket) -// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort -// Matched by: dport=wgPort -// -// 3. Ingress: Packets to WireGuard -// dst=127.0.0.1:wgPort -// Matched by: dport=wgPort -// -// 4. Ingress: Packets to proxy (after eBPF rewrite) -// dst=127.0.0.1:proxyPort -// Matched by: dport=proxyPort -// -// Rules are cleaned up when the firewall manager is closed. -func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error { - m.mutex.Lock() - defer m.mutex.Unlock() - - if !m.rawSupported { - return fmt.Errorf("raw table not available") - } - - wgPortStr := fmt.Sprintf("%d", wgPort) - proxyPortStr := fmt.Sprintf("%d", proxyPort) - - // Egress rules: match outgoing loopback UDP packets - outputRuleSport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--sport", wgPortStr, "-j", "NOTRACK"} - if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleSport...); err != nil { - return fmt.Errorf("add output sport notrack rule: %w", err) - } - - outputRuleDport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"} - if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleDport...); err != nil { - return fmt.Errorf("add output dport notrack rule: %w", err) - } - - // Ingress rules: match incoming loopback UDP packets - preroutingRuleWg := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"} - if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleWg...); err != nil { - return fmt.Errorf("add prerouting wg notrack rule: %w", err) - } - - preroutingRuleProxy := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", proxyPortStr, "-j", "NOTRACK"} - if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleProxy...); err != nil { - return fmt.Errorf("add prerouting proxy notrack rule: %w", err) - } - - log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort) - return nil -} - -func (m *Manager) initNoTrackChain() error { - if err := m.cleanupNoTrackChain(); err != nil { - log.Debugf("cleanup notrack chain: %v", err) - } - - if err := m.ipv4Client.NewChain(tableRaw, chainNameRaw); err != nil { - return fmt.Errorf("create chain: %w", err) - } - - jumpRule := []string{"-j", chainNameRaw} - - if err := m.ipv4Client.InsertUnique(tableRaw, chainOutput, 1, jumpRule...); err != nil { - if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil { - log.Debugf("delete orphan chain: %v", delErr) - } - return fmt.Errorf("add output jump rule: %w", err) - } - - if err := m.ipv4Client.InsertUnique(tableRaw, chainPrerouting, 1, jumpRule...); err != nil { - if delErr := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); delErr != nil { - log.Debugf("delete output jump rule: %v", delErr) - } - if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil { - log.Debugf("delete orphan chain: %v", delErr) - } - return fmt.Errorf("add prerouting jump rule: %w", err) - } - - m.rawSupported = true - return nil -} - -func (m *Manager) cleanupNoTrackChain() error { - exists, err := m.ipv4Client.ChainExists(tableRaw, chainNameRaw) - if err != nil { - if !m.rawSupported { - return nil - } - return fmt.Errorf("check chain exists: %w", err) - } - if !exists { - return nil - } - - jumpRule := []string{"-j", chainNameRaw} - - if err := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); err != nil { - return fmt.Errorf("remove output jump rule: %w", err) - } - - if err := m.ipv4Client.DeleteIfExists(tableRaw, chainPrerouting, jumpRule...); err != nil { - return fmt.Errorf("remove prerouting jump rule: %w", err) - } - - if err := m.ipv4Client.ClearAndDeleteChain(tableRaw, chainNameRaw); err != nil { - return fmt.Errorf("clear and delete chain: %w", err) - } - - m.rawSupported = false - return nil -} - func getConntrackEstablished() []string { return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"} } diff --git a/client/firewall/manager/firewall.go b/client/firewall/manager/firewall.go index 97a94d0f5..0eb376875 100644 --- a/client/firewall/manager/firewall.go +++ b/client/firewall/manager/firewall.go @@ -192,10 +192,6 @@ type Manager interface { // RemoveOutputDNAT removes an OUTPUT chain DNAT rule. RemoveOutputDNAT(localAddr netip.Addr, protocol Protocol, originalPort, translatedPort uint16) error - - // SetupEBPFProxyNoTrack creates static notrack rules for eBPF proxy loopback traffic. - // This prevents conntrack from interfering with WireGuard proxy communication. - SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error } // GenKey builds the rule id for this pair from the given format. diff --git a/client/firewall/nftables/manager_linux.go b/client/firewall/nftables/manager_linux.go index dbd5e4fa2..87651761f 100644 --- a/client/firewall/nftables/manager_linux.go +++ b/client/firewall/nftables/manager_linux.go @@ -12,7 +12,6 @@ import ( "github.com/google/nftables/expr" "github.com/hashicorp/go-multierror" log "github.com/sirupsen/logrus" - "golang.org/x/sys/unix" nberrors "github.com/netbirdio/netbird/client/errors" firewall "github.com/netbirdio/netbird/client/firewall/manager" @@ -55,9 +54,6 @@ type Manager struct { // IPv6 counterpart, nil when no v6 overlay. family6 *family - notrackOutputChain *nftables.Chain - notrackPreroutingChain *nftables.Chain - extMonitor *externalChainMonitor } @@ -170,10 +166,6 @@ func (m *Manager) initFirewall() (err error) { } } - if err := m.initNoTrackChains(workTable); err != nil { - log.Warnf("raw priority chains not available, notrack rules will be disabled: %v", err) - } - return nil } @@ -455,10 +447,6 @@ func (m *Manager) Flush() error { } } - if err := m.refreshNoTrackChains(); err != nil { - log.Errorf("failed to refresh notrack chains: %v", err) - } - return nil } @@ -571,176 +559,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort) } -const ( - chainNameRawOutput = "netbird-raw-out" - chainNameRawPrerouting = "netbird-raw-pre" -) - -// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic. -// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which -// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark). -// -// Traffic flows that need NOTRACK: -// -// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite) -// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort -// Matched by: sport=wgPort -// -// 2. Egress: Proxy -> WireGuard (via raw socket) -// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort -// Matched by: dport=wgPort -// -// 3. Ingress: Packets to WireGuard -// dst=127.0.0.1:wgPort -// Matched by: dport=wgPort -// -// 4. Ingress: Packets to proxy (after eBPF rewrite) -// dst=127.0.0.1:proxyPort -// Matched by: dport=proxyPort -// -// Rules are cleaned up when the firewall manager is closed. -func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error { - m.mutex.Lock() - defer m.mutex.Unlock() - - if m.notrackOutputChain == nil || m.notrackPreroutingChain == nil { - return fmt.Errorf("notrack chains not initialized") - } - - proxyPortBytes := binaryutil.BigEndian.PutUint16(proxyPort) - wgPortBytes := binaryutil.BigEndian.PutUint16(wgPort) - loopback := []byte{127, 0, 0, 1} - - // Egress rules: match outgoing loopback UDP packets - m.rConn.AddRule(&nftables.Rule{ - Table: m.notrackOutputChain.Table, - Chain: m.notrackOutputChain, - Exprs: []expr.Any{ - &expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 2}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // sport=wgPort - &expr.Counter{}, - &expr.Notrack{}, - }, - }) - m.rConn.AddRule(&nftables.Rule{ - Table: m.notrackOutputChain.Table, - Chain: m.notrackOutputChain, - Exprs: []expr.Any{ - &expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort - &expr.Counter{}, - &expr.Notrack{}, - }, - }) - - // Ingress rules: match incoming loopback UDP packets - m.rConn.AddRule(&nftables.Rule{ - Table: m.notrackPreroutingChain.Table, - Chain: m.notrackPreroutingChain, - Exprs: []expr.Any{ - &expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort - &expr.Counter{}, - &expr.Notrack{}, - }, - }) - m.rConn.AddRule(&nftables.Rule{ - Table: m.notrackPreroutingChain.Table, - Chain: m.notrackPreroutingChain, - Exprs: []expr.Any{ - &expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback}, - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}}, - &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, - &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: proxyPortBytes}, // dport=proxyPort - &expr.Counter{}, - &expr.Notrack{}, - }, - }) - - if err := m.rConn.Flush(); err != nil { - return fmt.Errorf("flush notrack rules: %w", err) - } - - log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort) - return nil -} - -func (m *Manager) initNoTrackChains(table *nftables.Table) error { - m.notrackOutputChain = m.rConn.AddChain(&nftables.Chain{ - Name: chainNameRawOutput, - Table: table, - Type: nftables.ChainTypeFilter, - Hooknum: nftables.ChainHookOutput, - Priority: nftables.ChainPriorityRaw, - }) - - m.notrackPreroutingChain = m.rConn.AddChain(&nftables.Chain{ - Name: chainNameRawPrerouting, - Table: table, - Type: nftables.ChainTypeFilter, - Hooknum: nftables.ChainHookPrerouting, - Priority: nftables.ChainPriorityRaw, - }) - - if err := m.rConn.Flush(); err != nil { - return fmt.Errorf("flush chain creation: %w", err) - } - - return nil -} - -func (m *Manager) refreshNoTrackChains() error { - chains, err := m.rConn.ListChainsOfTableFamily(nftables.TableFamilyIPv4) - if err != nil { - return fmt.Errorf("list chains: %w", err) - } - - tableName := getTableName() - for _, c := range chains { - if c.Table.Name != tableName { - continue - } - switch c.Name { - case chainNameRawOutput: - m.notrackOutputChain = c - case chainNameRawPrerouting: - m.notrackPreroutingChain = c - } - } - - return nil -} - func (m *Manager) createWorkTable() (*nftables.Table, error) { return m.createWorkTableFamily(nftables.TableFamilyIPv4) } diff --git a/client/firewall/nftables/routing_linux.go b/client/firewall/nftables/routing_linux.go index 4115c94bd..d619c5543 100644 --- a/client/firewall/nftables/routing_linux.go +++ b/client/firewall/nftables/routing_linux.go @@ -192,7 +192,7 @@ func (r *family) addPostroutingRules() { Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade), }, - // We need to exclude the loopback interface as this changes the ebpf proxy port + // We need to exclude the loopback interface as this changes the wg proxy port &expr.Meta{ Key: expr.MetaKeyOIFNAME, Register: 1, diff --git a/client/firewall/uspfilter/filter.go b/client/firewall/uspfilter/filter.go index 5e1366c1f..0c73400f1 100644 --- a/client/firewall/uspfilter/filter.go +++ b/client/firewall/uspfilter/filter.go @@ -879,12 +879,6 @@ func (m *Manager) resetState() { } } -// SetupEBPFProxyNoTrack is not supported by the userspace firewall: eBPF isn't -// used in userspace mode, so this should never be called. -func (m *Manager) SetupEBPFProxyNoTrack(uint16, uint16) error { - return errNotSupported -} - // UpdateSet updates the rule destinations associated with the given set // by merging the existing prefixes with the new ones, then deduplicating. func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error { diff --git a/client/firewall/uspfilter/interface_allower_windows.go b/client/firewall/uspfilter/interface_allower_windows.go index 7f525e28c..4cd0fe969 100644 --- a/client/firewall/uspfilter/interface_allower_windows.go +++ b/client/firewall/uspfilter/interface_allower_windows.go @@ -9,6 +9,7 @@ import ( log "github.com/sirupsen/logrus" nberrors "github.com/netbirdio/netbird/client/errors" + "github.com/netbirdio/netbird/client/internal/wincmd" ) type action string @@ -91,7 +92,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err if action == addRule { args = append(args, extraArgs...) } - netshCmd := GetSystem32Command("netsh") + netshCmd := wincmd.System32("netsh") cmd := exec.Command(netshCmd, args...) cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} return cmd.Run() @@ -100,7 +101,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err func isWindowsFirewallReachable() bool { args := []string{"advfirewall", "show", "allprofiles", "state"} - netshCmd := GetSystem32Command("netsh") + netshCmd := wincmd.System32("netsh") cmd := exec.Command(netshCmd, args...) cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} @@ -117,23 +118,10 @@ func isWindowsFirewallReachable() bool { func isFirewallRuleActive(ruleName string) bool { args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName} - netshCmd := GetSystem32Command("netsh") + netshCmd := wincmd.System32("netsh") cmd := exec.Command(netshCmd, args...) cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} _, err := cmd.Output() return err == nil } - -// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it -// in the path it will return the full path of a command assuming C:\windows\system32 as the base path. -func GetSystem32Command(command string) string { - _, err := exec.LookPath(command) - if err == nil { - return command - } - - log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command) - - return "C:\\windows\\system32\\" + command + ".exe" -} diff --git a/client/iface/configurer/allowedips.go b/client/iface/configurer/allowedips.go new file mode 100644 index 000000000..193197d4a --- /dev/null +++ b/client/iface/configurer/allowedips.go @@ -0,0 +1,226 @@ +package configurer + +import ( + "net" + "net/netip" + "slices" + "sync" + + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// allowedIPStore mirrors the allowed IPs configured on each peer of a device. +// +// A configurer is the only writer of its device's peer set, so the mirror is authoritative +// by construction. It spares the paths that have to rewrite one peer's allowed IPs a full +// device dump just to recover prefixes the process already configured itself. +// +// An allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away +// from whichever peer held it before, and the configurer leaves that handover to the device +// rather than removing the prefix from the previous holder itself. The store tracks the +// owner of each prefix and performs the same handover, so rewriting one peer's list never +// takes a prefix back from the peer that owns it now. +// +// Its own lock guards the map alone, not the device write it accompanies. Consistency +// between the two rests on the caller serializing every configurer call, which WGIface +// does with its mutex; two unserialized writers would interleave a device write with the +// record of a different one. +// +// An operator reconfiguring the device out of band, through `wg set` or the UAPI socket, +// is the one way the mirror can still go stale. A peer missing from it falls back to the +// device, which reseats that peer's prefixes and their ownership; a peer that is present +// does not, so one recorded from empty while the device already held prefixes keeps only +// what was recorded, and the next endpoint removal drops the rest. +type allowedIPStore struct { + mu sync.RWMutex + peers map[wgtypes.Key][]netip.Prefix + owners map[netip.Prefix]wgtypes.Key +} + +func newAllowedIPStore() *allowedIPStore { + return &allowedIPStore{ + peers: make(map[wgtypes.Key][]netip.Prefix), + owners: make(map[netip.Prefix]wgtypes.Key), + } +} + +// get returns the prefixes recorded for a peer, and whether the peer is known at all. +// The caller receives a copy and may retain or modify it freely. +func (s *allowedIPStore) get(key wgtypes.Key) ([]netip.Prefix, bool) { + s.mu.RLock() + defer s.mu.RUnlock() + + prefixes, ok := s.peers[key] + if !ok { + return nil, false + } + return slices.Clone(prefixes), true +} + +// set replaces the prefixes recorded for a peer. +func (s *allowedIPStore) set(key wgtypes.Key, prefixes []netip.Prefix) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + s.releaseLocked(k) + + normalized := normalizePrefixes(prefixes) + for _, prefix := range normalized { + s.claimLocked(k, prefix) + } + s.peers[k] = normalized +} + +// add records prefixes on a peer without dropping the ones already there, matching the +// union semantics of a peer update that does not replace its allowed IPs. It records the +// peer if it is not known yet, so it belongs to the operations that create a peer on the +// device rather than to the update-only ones. +func (s *allowedIPStore) add(key wgtypes.Key, prefixes []netip.Prefix) { + s.mu.Lock() + defer s.mu.Unlock() + + s.mergeLocked(key, prefixes) +} + +// addExisting is add for an update-only device operation. Such an operation is a silent +// no-op when the peer is absent, so recording a peer here would leave the store claiming +// prefixes the device never took, and the peer would then be recreated by the next endpoint +// removal, stealing those allowed IPs from the peer that legitimately holds them. +func (s *allowedIPStore) addExisting(key wgtypes.Key, prefixes []netip.Prefix) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + if _, ok := s.peers[k]; !ok { + return + } + s.mergeLocked(k, prefixes) +} + +// ensure records a peer with no prefixes unless it is already known. A device operation +// that is not update-only creates the peer when it is absent, so it has to be recorded even +// when it configures nothing else; otherwise the peer exists on the device while the store +// treats it as unknown, and a prefix later handed over to it is not accounted for. +func (s *allowedIPStore) ensure(key wgtypes.Key) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + if _, ok := s.peers[k]; !ok { + s.peers[k] = nil + } +} + +// forget drops every prefix recorded for a peer. +func (s *allowedIPStore) forget(key wgtypes.Key) { + s.mu.Lock() + defer s.mu.Unlock() + + k := key + s.releaseLocked(k) + delete(s.peers, k) +} + +// reset drops every peer, mirroring a device reconfiguration that replaces the peer set. +func (s *allowedIPStore) reset() { + s.mu.Lock() + defer s.mu.Unlock() + + s.peers = make(map[wgtypes.Key][]netip.Prefix) + s.owners = make(map[netip.Prefix]wgtypes.Key) +} + +// mergeLocked unions normalized prefixes into a peer and transfers their ownership. +// The caller must hold s.mu for writing. +func (s *allowedIPStore) mergeLocked(k wgtypes.Key, prefixes []netip.Prefix) { + merged := s.peers[k] + for _, prefix := range prefixes { + prefix = normalizePrefix(prefix) + s.claimLocked(k, prefix) + if !slices.Contains(merged, prefix) { + merged = append(merged, prefix) + } + } + s.peers[k] = merged +} + +// claimLocked hands a prefix over to a peer, taking it from its previous owner the way the +// device does when the same prefix is configured on a second peer. +func (s *allowedIPStore) claimLocked(k wgtypes.Key, prefix netip.Prefix) { + if owner, ok := s.owners[prefix]; ok && owner != k { + s.peers[owner] = slices.DeleteFunc(s.peers[owner], func(p netip.Prefix) bool { + return p == prefix + }) + } + s.owners[prefix] = k +} + +// releaseLocked drops a peer's claim on every prefix it currently holds. +func (s *allowedIPStore) releaseLocked(k wgtypes.Key) { + for _, prefix := range s.peers[k] { + if s.owners[prefix] == k { + delete(s.owners, prefix) + } + } +} + +// normalizePrefix puts a prefix into the form the store recognises it by. It clears the +// host bits, which a device does on its own, so a caller passing 10.20.0.1/16 still matches +// the 10.20.0.0/16 read back from the device; and it unmaps a v4-mapped prefix so that it +// compares equal to, and marshals like, the plain v4 prefix for the same network. +// +// Masking comes first because it also decides the address family: only a prefix at least 96 +// bits long keeps the mapped marker through the mask, so a shorter prefix inside the mapped +// range is a genuine v6 prefix and unmapping it would yield an invalid v4 prefix. +func normalizePrefix(prefix netip.Prefix) netip.Prefix { + masked := prefix.Masked() + + addr := masked.Addr() + if !addr.Is4In6() { + return masked + } + return netip.PrefixFrom(addr.Unmap(), masked.Bits()-96) +} + +// normalizePrefixes returns a normalized copy without changing the caller's slice. +func normalizePrefixes(prefixes []netip.Prefix) []netip.Prefix { + normalized := make([]netip.Prefix, len(prefixes)) + for i, prefix := range prefixes { + normalized[i] = normalizePrefix(prefix) + } + return normalized +} + +// ipNetsToPrefixes converts addresses read back from a device. Unmap keeps a v4-mapped v6 +// address comparable to the plain v4 prefix the configurer was given. +func ipNetsToPrefixes(ipNets []net.IPNet) []netip.Prefix { + prefixes := make([]netip.Prefix, 0, len(ipNets)) + for _, ipNet := range ipNets { + addr, ok := netip.AddrFromSlice(ipNet.IP) + if !ok { + continue + } + + ones, maskBits := ipNet.Mask.Size() + // A device may report a v4 prefix as a v4-mapped address. Align the address form with + // the mask rather than unmapping on sight: a 32 bit mask always describes v4, while a + // 128 bit mask describes v4 only when it covers the mapped prefix, so a genuine v6 + // prefix inside the mapped range stays v6 instead of being dropped as invalid. + if addr.Is4In6() { + switch { + case maskBits == 32: + addr = addr.Unmap() + case maskBits == 128 && ones >= 96: + addr, ones = addr.Unmap(), ones-96 + } + } + + prefix := netip.PrefixFrom(addr, ones) + if !prefix.IsValid() { + continue + } + prefixes = append(prefixes, prefix.Masked()) + } + return prefixes +} diff --git a/client/iface/configurer/allowedips_test.go b/client/iface/configurer/allowedips_test.go new file mode 100644 index 000000000..1d272d8c4 --- /dev/null +++ b/client/iface/configurer/allowedips_test.go @@ -0,0 +1,263 @@ +package configurer + +import ( + "net" + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// The store keys on the parsed key, so the tests use two distinct ones rather than names. +var ( + testPeer = wgtypes.Key{1} + otherPeer = wgtypes.Key{2} +) + +func TestAllowedIPStoreUnknownPeer(t *testing.T) { + s := newAllowedIPStore() + + prefixes, ok := s.get(testPeer) + assert.False(t, ok, "an unconfigured peer must be reported as unknown, not as one without prefixes") + assert.Nil(t, prefixes, "an unknown peer has no prefixes") +} + +func TestAllowedIPStoreAddUnions(t *testing.T) { + s := newAllowedIPStore() + overlay := netip.MustParsePrefix("100.64.0.1/32") + routed := netip.MustParsePrefix("10.20.0.0/16") + + s.set(testPeer, []netip.Prefix{overlay}) + // A peer update does not replace allowed IPs, and a repeated prefix must not be doubled. + s.add(testPeer, []netip.Prefix{overlay, routed}) + + prefixes, ok := s.get(testPeer) + require.True(t, ok, "peer must be known after set") + assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "add must union rather than replace") +} + +func TestAllowedIPStoreGetReturnsCopy(t *testing.T) { + s := newAllowedIPStore() + overlay := netip.MustParsePrefix("100.64.0.1/32") + s.set(testPeer, []netip.Prefix{overlay}) + + prefixes, ok := s.get(testPeer) + require.True(t, ok, "peer must be known after set") + prefixes[0] = netip.MustParsePrefix("0.0.0.0/0") + + stored, _ := s.get(testPeer) + assert.Equal(t, []netip.Prefix{overlay}, stored, "a caller mutating the returned slice must not corrupt the store") +} + +func TestAllowedIPStoreForgetAndReset(t *testing.T) { + s := newAllowedIPStore() + s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")}) + s.set(otherPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")}) + + s.forget(testPeer) + _, ok := s.get(testPeer) + assert.False(t, ok, "a forgotten peer must be unknown") + _, ok = s.get(otherPeer) + assert.True(t, ok, "forgetting one peer must not touch the others") + + s.reset() + _, ok = s.get(otherPeer) + assert.False(t, ok, "reset must drop every peer") +} + +func TestIPNetsToPrefixes(t *testing.T) { + tests := []struct { + name string + ipNet net.IPNet + want string + }{ + { + name: "v4", + ipNet: net.IPNet{IP: net.IP{10, 20, 0, 0}, Mask: net.CIDRMask(16, 32)}, + want: "10.20.0.0/16", + }, + { + name: "v4 mapped under a 128 bit mask", + ipNet: net.IPNet{IP: net.ParseIP("10.20.0.0"), Mask: net.CIDRMask(112, 128)}, + want: "10.20.0.0/16", + }, + { + name: "v6", + ipNet: net.IPNet{IP: net.ParseIP("fd00::"), Mask: net.CIDRMask(64, 128)}, + want: "fd00::/64", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := ipNetsToPrefixes([]net.IPNet{tc.ipNet}) + require.Len(t, got, 1, "the address must be converted, not dropped") + assert.Equal(t, tc.want, got[0].String(), "converted prefix") + }) + } +} + +func TestIPNetsToPrefixesRoundTrip(t *testing.T) { + prefixes := []netip.Prefix{ + netip.MustParsePrefix("100.64.0.1/32"), + netip.MustParsePrefix("10.20.0.0/16"), + netip.MustParsePrefix("fd00::/64"), + } + + assert.Equal(t, prefixes, ipNetsToPrefixes(prefixesToIPNets(prefixes)), + "prefixes handed to a device must come back unchanged") +} + +func TestAllowedIPStoreNormalizesMappedPrefixes(t *testing.T) { + s := newAllowedIPStore() + v4 := netip.MustParsePrefix("10.20.0.0/16") + mapped := netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112) + + s.set(testPeer, []netip.Prefix{mapped}) + // A v4 rule only matches a v4-mapped address once it has been unmapped, so the store must + // hold the plain form and recognise the two spellings as the same prefix. + s.add(testPeer, []netip.Prefix{v4}) + + prefixes, ok := s.get(testPeer) + require.True(t, ok, "peer must be known after set") + assert.Equal(t, []netip.Prefix{v4}, prefixes, "a mapped prefix must be stored unmapped and not duplicated") +} + +func TestNormalizePrefix(t *testing.T) { + v4 := netip.MustParsePrefix("10.20.0.0/16") + v6 := netip.MustParsePrefix("fd00::/64") + + assert.Equal(t, v4, normalizePrefix(v4), "a plain v4 prefix is unchanged") + assert.Equal(t, v6, normalizePrefix(v6), "a real v6 prefix is unchanged") + assert.Equal(t, v4, normalizePrefix(netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)), + "a mapped prefix under a 128 bit mask becomes plain v4") + // A prefix shorter than /96 inside the mapped range is a genuine v6 prefix. Unmapping it + // would pair a v4 address with a v6 sized mask, which is invalid, and the store would then + // record a zero prefix that can never recreate the allowed IP. + for _, tc := range []string{"::ffff:0:0/64", "::ffff:1.2.3.4/80", "::ffff:1.2.3.4/95"} { + got := normalizePrefix(netip.MustParsePrefix(tc)) + assert.True(t, got.IsValid(), "%s must normalize to a valid prefix", tc) + assert.False(t, got.Addr().Is4(), "%s must stay v6", tc) + } +} + +func TestAllowedIPStoreAddExistingDoesNotCreate(t *testing.T) { + s := newAllowedIPStore() + routed := netip.MustParsePrefix("10.20.0.0/16") + + // An update-only device operation on an absent peer is a silent no-op, so nothing may be + // recorded for a peer the store does not already know. + s.addExisting(testPeer, []netip.Prefix{routed}) + _, ok := s.get(testPeer) + assert.False(t, ok, "addExisting must not record an unknown peer") + + overlay := netip.MustParsePrefix("100.64.0.1/32") + s.set(testPeer, []netip.Prefix{overlay}) + s.addExisting(testPeer, []netip.Prefix{routed}) + + prefixes, _ := s.get(testPeer) + assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "addExisting must union onto a known peer") +} + +func TestAllowedIPStoreHandsPrefixOverToTheNewOwner(t *testing.T) { + s := newAllowedIPStore() + routed := netip.MustParsePrefix("10.20.0.0/16") + other := otherPeer + + s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32"), routed}) + s.set(other, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")}) + + // The device takes an allowed IP away from its previous holder when it is configured on + // another peer, so the store must do the same rather than list it under both. + s.addExisting(other, []netip.Prefix{routed}) + + previous, _ := s.get(testPeer) + assert.NotContains(t, previous, routed, "the previous owner must lose the prefix") + current, _ := s.get(other) + assert.Contains(t, current, routed, "the new owner must hold the prefix") +} + +func TestAllowedIPStoreForgetReleasesOwnership(t *testing.T) { + s := newAllowedIPStore() + routed := netip.MustParsePrefix("10.20.0.0/16") + + s.set(testPeer, []netip.Prefix{routed}) + s.forget(testPeer) + s.set(otherPeer, []netip.Prefix{routed}) + + // A forgotten peer must not be resurrected as a key in the peer map by a later claim. + _, ok := s.get(testPeer) + assert.False(t, ok, "the forgotten peer must stay unknown") + current, _ := s.get(otherPeer) + assert.Equal(t, []netip.Prefix{routed}, current, "the new owner must hold the prefix") +} + +func TestNormalizePrefixClearsHostBits(t *testing.T) { + // A device stores a prefix masked, so a caller passing host bits must still match what a + // device fallback seeded, otherwise that prefix could never be removed by value. + assert.Equal(t, netip.MustParsePrefix("10.20.0.0/16"), + normalizePrefix(netip.MustParsePrefix("10.20.0.1/16")), "host bits must be cleared") + assert.Equal(t, netip.MustParsePrefix("fd00::/64"), + normalizePrefix(netip.MustParsePrefix("fd00::1/64")), "host bits must be cleared for v6") +} + +func TestIPNetsToPrefixesKeepsV6InTheMappedRange(t *testing.T) { + // ::ffff:0:0/64 reads as v4-mapped but is a genuine v6 prefix: unmapping it would leave a + // v4 address under a 64 bit mask, which is invalid, and the allowed IP would be dropped. + got := ipNetsToPrefixes([]net.IPNet{{ + IP: net.ParseIP("::ffff:0:0"), + Mask: net.CIDRMask(64, 128), + }}) + + require.Len(t, got, 1, "the prefix must be converted, not dropped") + assert.False(t, got[0].Addr().Is4(), "a v6 prefix in the mapped range must not become v4") + assert.Equal(t, 64, got[0].Bits(), "the prefix length must survive the conversion") +} + +func TestPrefixesToIPNetsNormalizes(t *testing.T) { + // net.IPNet prints a v4-mapped address as v4 but takes the length from its 16 byte + // mask, so an unnormalized ::ffff:10.1.2.3/64 reaches a userspace device as 10.1.2.3/0, + // an allowed IP that matches every v4 address. + tests := []struct { + name string + given string + want string + }{ + {name: "mapped below /96", given: "::ffff:10.1.2.3/64", want: "::/64"}, + {name: "mapped at /112", given: "::ffff:10.1.2.3/112", want: "10.1.0.0/16"}, + {name: "host bits are cleared", given: "10.20.0.1/16", want: "10.20.0.0/16"}, + {name: "v6 is untouched", given: "fd00::1/64", want: "fd00::/64"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := prefixesToIPNets([]netip.Prefix{netip.MustParsePrefix(tc.given)}) + require.Len(t, got, 1, "the prefix must be converted, not dropped") + assert.Equal(t, tc.want, got[0].String(), "what the device is given") + assert.NotEqual(t, 0, mustOnes(t, got[0]), "a device must never be given a zero length allowed IP") + }) + } +} + +func mustOnes(t *testing.T, ipNet net.IPNet) int { + t.Helper() + + ones, _ := ipNet.Mask.Size() + return ones +} + +// TestPrefixesToIPNetsAgreesWithTheStore pins the property the store depends on: what a +// device is given and what is recorded for it are the same prefix. +func TestPrefixesToIPNetsAgreesWithTheStore(t *testing.T) { + for _, given := range []string{"::ffff:10.1.2.3/64", "::ffff:10.1.2.3/112", "10.20.0.1/16", "fd00::1/64"} { + prefix := netip.MustParsePrefix(given) + + toDevice := prefixesToIPNets([]netip.Prefix{prefix}) + recorded := normalizePrefix(prefix) + + assert.Equal(t, recorded.String(), toDevice[0].String(), + "%s must reach the device in the form the store records", given) + } +} diff --git a/client/iface/configurer/common.go b/client/iface/configurer/common.go index 10162d703..40f8209e9 100644 --- a/client/iface/configurer/common.go +++ b/client/iface/configurer/common.go @@ -19,12 +19,18 @@ func buildPresharedKeyConfig(peerKey wgtypes.Key, psk wgtypes.Key, updateOnly bo } } +// prefixesToIPNets converts prefixes on their way to a device. It is the only place that +// conversion happens, so it also normalizes: the device is then given the same form the +// store records, and a v4-mapped prefix cannot reach net.IPNet, which prints such an +// address as v4 while taking the length from its 16 byte mask and so turns +// ::ffff:10.1.2.3/64 into 10.1.2.3/0 — an allowed IP matching every v4 address. func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet { ipNets := make([]net.IPNet, len(prefixes)) for i, prefix := range prefixes { + normalized := normalizePrefix(prefix) ipNets[i] = net.IPNet{ - IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP - Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask + IP: normalized.Addr().AsSlice(), + Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()), } } return ipNets diff --git a/client/iface/configurer/kernel_unix.go b/client/iface/configurer/kernel_unix.go index da69c2a35..3a95249c1 100644 --- a/client/iface/configurer/kernel_unix.go +++ b/client/iface/configurer/kernel_unix.go @@ -6,6 +6,7 @@ import ( "fmt" "net" "net/netip" + "slices" "time" log "github.com/sirupsen/logrus" @@ -18,16 +19,22 @@ import ( type KernelConfigurer struct { deviceName string statsCache *statsCache + allowedIPs *allowedIPStore } +// NewKernelConfigurer creates a configurer with an empty allowed IP mirror +// and a statistics cache for the named kernel device. func NewKernelConfigurer(deviceName string) *KernelConfigurer { c := &KernelConfigurer{ deviceName: deviceName, + allowedIPs: newAllowedIPStore(), } c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats) return c } +// ConfigureInterface sets the device key, port and firewall mark, replacing all peers. +// The allowed IP mirror is reset only after the device accepts the configuration. func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error { log.Debugf("adding Wireguard private key") key, err := wgtypes.ParseKey(privateKey) @@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error if err != nil { return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port) } + + c.allowedIPs.reset() return nil } @@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda } cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly) - return c.configure(cfg) + if err := c.configure(cfg); err != nil { + return err + } + + // Without updateOnly this creates the peer when it is absent, so the store has to + // know about it even though no allowed IP was configured. + if !updateOnly { + c.allowedIPs.ensure(parsedPeerKey) + } + return nil } +// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set. +// Prefixes assigned to this peer are transferred from their previous owners. func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { @@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, if err != nil { return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String()) } + + c.allowedIPs.add(peerKeyParsed, allowedIps) return nil } +// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured. +// Neither the netlink API nor the userspace one can clear an endpoint in place, so the peer +// is removed and re-added with the allowed IPs it already had. func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err } - // Get the existing peer to preserve its allowed IPs - existingPeer, err := c.getPeer(c.deviceName, peerKey) + allowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { - return fmt.Errorf("get peer: %w", err) + return err } removePeerCfg := wgtypes.PeerConfig{ @@ -104,26 +128,27 @@ func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error { } if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil { - return fmt.Errorf(`error removing peer %s from interface %s: %w`, peerKey, c.deviceName, err) + return fmt.Errorf("remove peer %s from interface %s: %w", peerKey, c.deviceName, err) } - //Re-add the peer without the endpoint but same AllowedIPs reAddPeerCfg := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, - AllowedIPs: existingPeer.AllowedIPs, + AllowedIPs: prefixesToIPNets(allowedIPs), ReplaceAllowedIPs: true, } if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil { + c.allowedIPs.forget(peerKeyParsed) return fmt.Errorf( - `error re-adding peer %s to interface %s with allowed IPs %v: %w`, - peerKey, c.deviceName, existingPeer.AllowedIPs, err, + "re-add peer %s to interface %s with allowed IPs %v: %w", + peerKey, c.deviceName, allowedIPs, err, ) } return nil } +// RemovePeer removes a peer and forgets its allowed IPs after a successful device write. func (c *KernelConfigurer) RemovePeer(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { @@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error { if err != nil { return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName) } + + c.allowedIPs.forget(peerKeyParsed) return nil } +// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op. func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipNet := net.IPNet{ - IP: allowedIP.Addr().AsSlice(), - Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), - } - peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err @@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: false, - AllowedIPs: []net.IPNet{ipNet}, + AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}), } config := wgtypes.Config{ @@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) if err != nil { return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP) } + + c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP}) return nil } +// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs. +// A prefix not assigned to the peer is a no-op. func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipNet := net.IPNet{ - IP: allowedIP.Addr().AsSlice(), - Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), - } - peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return fmt.Errorf("parse peer key: %w", err) } - existingPeer, err := c.getPeer(c.deviceName, peerKey) + currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { - return fmt.Errorf("get peer: %w", err) + return err } - newAllowedIPs := existingPeer.AllowedIPs - - for i, existingAllowedIP := range existingPeer.AllowedIPs { - if existingAllowedIP.String() == ipNet.String() { - newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic - break - } + idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP)) + if idx < 0 { + return nil } + newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1) peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: true, - AllowedIPs: newAllowedIPs, + AllowedIPs: prefixesToIPNets(newAllowedIPs), } config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - err = c.configure(config) - if err != nil { + if err := c.configure(config); err != nil { return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err) } + + c.allowedIPs.set(peerKeyParsed, newAllowedIPs) return nil } -func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) { +// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device +// only for a peer the store has not seen. Dumping the device costs a netlink round trip +// proportional to the whole network map, and this runs on every relay and ICE transition. +func (c *KernelConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) { + if prefixes, ok := c.allowedIPs.get(peerKey); ok { + return prefixes, nil + } + + existingPeer, err := c.getPeer(c.deviceName, peerKey) + if err != nil { + return nil, fmt.Errorf("get peer: %w", err) + } + + prefixes := ipNetsToPrefixes(existingPeer.AllowedIPs) + c.allowedIPs.set(peerKey, prefixes) + return prefixes, nil +} + +// getPeer scans the device for one peer. wgtypes.Key is an array, so the comparison is a +// plain equality: Key.String would base64 encode into a fresh allocation for every peer. +func (c *KernelConfigurer) getPeer(ifaceName string, peerPubKey wgtypes.Key) (wgtypes.Peer, error) { wg, err := wgctrl.New() if err != nil { return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err) @@ -231,7 +271,7 @@ func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err) } for _, peer := range wgDevice.Peers { - if peer.PublicKey.String() == peerPubKey { + if peer.PublicKey == peerPubKey { return peer, nil } } diff --git a/client/iface/configurer/usp.go b/client/iface/configurer/usp.go index 2be1b861e..334d99369 100644 --- a/client/iface/configurer/usp.go +++ b/client/iface/configurer/usp.go @@ -8,6 +8,7 @@ import ( "net/netip" "os" "runtime" + "slices" "strconv" "strings" "time" @@ -41,31 +42,38 @@ type WGUSPConfigurer struct { deviceName string activityRecorder *bind.ActivityRecorder statsCache *statsCache + allowedIPs *allowedIPStore uapiListener net.Listener } +// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener. func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer { wgCfg := &WGUSPConfigurer{ device: device, deviceName: deviceName, activityRecorder: activityRecorder, + allowedIPs: newAllowedIPStore(), } wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats) wgCfg.startUAPI() return wgCfg } +// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener. func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer { wgCfg := &WGUSPConfigurer{ device: device, deviceName: deviceName, activityRecorder: activityRecorder, + allowedIPs: newAllowedIPStore(), } wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats) return wgCfg } +// ConfigureInterface sets the device key, port and firewall mark, replacing all peers. +// The allowed IP mirror is reset only after the device accepts the configuration. func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error { log.Debugf("adding Wireguard private key") key, err := wgtypes.ParseKey(privateKey) @@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error ListenPort: &port, } - return c.device.IpcSet(toWgUserspaceString(config)) + if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { + return err + } + + c.allowedIPs.reset() + return nil } // SetPresharedKey sets the preshared key for a peer. @@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat } cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly) - return c.device.IpcSet(toWgUserspaceString(cfg)) + if err := c.device.IpcSet(toWgUserspaceString(cfg)); err != nil { + return err + } + + // Without updateOnly this creates the peer when it is absent, so the store has to + // know about it even though no allowed IP was configured. + if !updateOnly { + c.allowedIPs.ensure(parsedPeerKey) + } + return nil } +// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set. +// It validates the endpoint before writing and records changes after a successful write. func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err } + + // Everything that can fail is done before the device is touched, so a failure here + // cannot leave the device holding a peer that the activity recorder and the allowed + // IP store never learned about. + var addrPort netip.AddrPort + if endpoint != nil { + addr, err := netip.ParseAddr(endpoint.IP.String()) + if err != nil { + return fmt.Errorf("parse endpoint address: %w", err) + } + addrPort = netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port)) + } + peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, ReplaceAllowedIPs: false, @@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, } if endpoint != nil { - addr, err := netip.ParseAddr(endpoint.IP.String()) - if err != nil { - return fmt.Errorf("failed to parse endpoint address: %w", err) - } - addrPort := netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port)) c.activityRecorder.UpsertAddress(peerKey, addrPort) } + + c.allowedIPs.add(peerKeyParsed, allowedIps) return nil } +// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured. +// The UAPI cannot clear an endpoint in place, so the peer is removed and re-added with the +// allowed IPs it already had. func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return fmt.Errorf("parse peer key: %w", err) } - ipcStr, err := c.device.IpcGet() + allowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { - return fmt.Errorf("get IPC config: %w", err) + return err } - // Parse current status to get allowed IPs for the peer - stats, err := parseStatus(c.deviceName, ipcStr) - if err != nil { - return fmt.Errorf("parse IPC config: %w", err) - } - - var allowedIPs []net.IPNet - found := false - for _, peer := range stats.Peers { - if peer.PublicKey == peerKey { - allowedIPs = peer.AllowedIPs - found = true - break - } - } - if !found { - return fmt.Errorf("peer %s not found", peerKey) - } - - // remove the peer from the WireGuard configuration peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, Remove: true, @@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { Peers: []wgtypes.PeerConfig{peer}, } if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil { - return fmt.Errorf("failed to remove peer: %s", ipcErr) + return fmt.Errorf("remove peer: %w", ipcErr) } - // Build the peer config peer = wgtypes.PeerConfig{ PublicKey: peerKeyParsed, ReplaceAllowedIPs: true, - AllowedIPs: allowedIPs, + AllowedIPs: prefixesToIPNets(allowedIPs), } config = wgtypes.Config{ @@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error { } if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { - return fmt.Errorf("remove endpoint address: %w", err) + c.allowedIPs.forget(peerKeyParsed) + return fmt.Errorf("re-add peer without endpoint: %w", err) } return nil } +// RemovePeer removes a peer, then clears its activity and allowed IP records. +// A failed device write leaves both records intact. func (c *WGUSPConfigurer) RemovePeer(peerKey string) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { @@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error { config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - ipcErr := c.device.IpcSet(toWgUserspaceString(config)) - - c.activityRecorder.Remove(peerKey) - return ipcErr -} - -func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipNet := net.IPNet{ - IP: allowedIP.Addr().AsSlice(), - Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()), + if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil { + return ipcErr } + c.activityRecorder.Remove(peerKey) + c.allowedIPs.forget(peerKeyParsed) + return nil +} + +// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op. +func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { peerKeyParsed, err := wgtypes.ParseKey(peerKey) if err != nil { return err @@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: false, - AllowedIPs: []net.IPNet{ipNet}, + AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}), } config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - return c.device.IpcSet(toWgUserspaceString(config)) + if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { + return err + } + + c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP}) + return nil } +// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs. +// It returns ErrAllowedIPNotFound if the prefix is not assigned to the peer. func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error { - ipc, err := c.device.IpcGet() - if err != nil { - return err - } - peerKeyParsed, err := wgtypes.ParseKey(peerKey) + if err != nil { + return fmt.Errorf("parse peer key: %w", err) + } + + currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed) if err != nil { return err } - hexKey := hex.EncodeToString(peerKeyParsed[:]) - lines := strings.Split(ipc, "\n") + idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP)) + if idx < 0 { + return ErrAllowedIPNotFound + } + newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1) peer := wgtypes.PeerConfig{ PublicKey: peerKeyParsed, UpdateOnly: true, ReplaceAllowedIPs: true, - AllowedIPs: []net.IPNet{}, + AllowedIPs: prefixesToIPNets(newAllowedIPs), } - foundPeer := false - removedAllowedIP := false - ip := allowedIP.String() - - for _, line := range lines { - line = strings.TrimSpace(line) - - // If we're within the details of the found peer and encounter another public key, - // this means we're starting another peer's details. So, reset the flag. - if strings.HasPrefix(line, "public_key=") && foundPeer { - foundPeer = false - } - - // Identify the peer with the specific public key - if line == fmt.Sprintf("public_key=%s", hexKey) { - foundPeer = true - } - - // If we're within the details of the found peer and find the specific allowed IP, skip this line - if foundPeer && line == "allowed_ip="+ip { - removedAllowedIP = true - continue - } - - // Append the line to the output string - if foundPeer && strings.HasPrefix(line, "allowed_ip=") { - allowedIPStr := strings.TrimPrefix(line, "allowed_ip=") - _, ipNet, err := net.ParseCIDR(allowedIPStr) - if err != nil { - return err - } - peer.AllowedIPs = append(peer.AllowedIPs, *ipNet) - } - } - - if !removedAllowedIP { - return ErrAllowedIPNotFound - } config := wgtypes.Config{ Peers: []wgtypes.PeerConfig{peer}, } - return c.device.IpcSet(toWgUserspaceString(config)) + if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil { + return fmt.Errorf("remove allowed IP %s: %w", allowedIP, err) + } + + c.allowedIPs.set(peerKeyParsed, newAllowedIPs) + return nil +} + +// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device +// only for a peer the store has not seen. Reading them back means dumping and parsing the +// whole device configuration, and this runs on every relay and ICE transition. +func (c *WGUSPConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) { + if prefixes, ok := c.allowedIPs.get(peerKey); ok { + return prefixes, nil + } + + ipcStr, err := c.device.IpcGet() + if err != nil { + return nil, fmt.Errorf("get IPC config: %w", err) + } + + stats, err := parseStatus(c.deviceName, ipcStr) + if err != nil { + return nil, fmt.Errorf("parse IPC config: %w", err) + } + + // parseStatus reports keys in their textual form, so the comparison needs it once. + wanted := peerKey.String() + for _, peer := range stats.Peers { + if peer.PublicKey != wanted { + continue + } + + prefixes := ipNetsToPrefixes(peer.AllowedIPs) + c.allowedIPs.set(peerKey, prefixes) + return prefixes, nil + } + + return nil, ErrPeerNotFound } func (c *WGUSPConfigurer) FullStats() (*Stats, error) { diff --git a/client/iface/configurer/usp_allowedips_test.go b/client/iface/configurer/usp_allowedips_test.go new file mode 100644 index 000000000..fba0ca546 --- /dev/null +++ b/client/iface/configurer/usp_allowedips_test.go @@ -0,0 +1,318 @@ +package configurer + +import ( + "net" + "net/netip" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + wgconn "golang.zx2c4.com/wireguard/conn" + wgdevice "golang.zx2c4.com/wireguard/device" + "golang.zx2c4.com/wireguard/tun/tuntest" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" + + "github.com/netbirdio/netbird/client/iface/bind" +) + +// newTestUSPConfigurer builds a configurer over a real wireguard-go device backed by an +// in-memory TUN. The device stays down, so no socket is opened and no privileges are needed. +func newTestUSPConfigurer(t *testing.T) *WGUSPConfigurer { + t.Helper() + + tun := tuntest.NewChannelTUN() + dev := wgdevice.NewDevice(tun.TUN(), wgconn.NewDefaultBind(), wgdevice.NewLogger(wgdevice.LogLevelSilent, "")) + t.Cleanup(dev.Close) + + c := NewUSPConfigurerNoUAPI(dev, "wgtest0", bind.NewActivityRecorder()) + + key, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate device private key") + require.NoError(t, c.ConfigureInterface(key.String(), 0), "configure test device") + + return c +} + +// seedPeers adds count peers, each with a /32 overlay address, and returns their public keys. +func seedPeers(t *testing.T, c *WGUSPConfigurer, count int) []string { + t.Helper() + + keys := make([]string, 0, count) + for i := 0; i < count; i++ { + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + pub := priv.PublicKey().String() + + addr := netip.PrefixFrom(netip.AddrFrom4([4]byte{100, 64, byte(i >> 8), byte(i)}), 32) + require.NoError(t, c.UpdatePeer(pub, []netip.Prefix{addr}, 25*time.Second, nil, nil), "add peer") + keys = append(keys, pub) + } + return keys +} + +func peerAllowedIPs(t *testing.T, c *WGUSPConfigurer, peerKey string) []string { + t.Helper() + + stats, err := c.FullStats() + require.NoError(t, err, "read device stats") + + for _, p := range stats.Peers { + if p.PublicKey != peerKey { + continue + } + got := make([]string, 0, len(p.AllowedIPs)) + for _, ipNet := range p.AllowedIPs { + got = append(got, ipNet.String()) + } + return got + } + t.Fatalf("peer %s not found on device", peerKey) + return nil +} + +// TestRemoveEndpointAddressPreservesRoutedAllowedIPs covers the prefixes the route manager +// attaches to a routing peer through AddAllowedIP. Those are not known to the peer.Conn that +// triggers the endpoint removal, so dropping them here would silently blackhole every route +// behind that peer on each relay or ICE disconnect. +func TestRemoveEndpointAddressPreservesRoutedAllowedIPs(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 3)[1] + + routed := []netip.Prefix{ + netip.MustParsePrefix("10.20.0.0/16"), + netip.MustParsePrefix("192.168.7.0/24"), + } + for _, prefix := range routed { + require.NoError(t, c.AddAllowedIP(peerKey, prefix), "add routed prefix") + } + + before := peerAllowedIPs(t, c, peerKey) + require.Len(t, before, 3, "peer should hold its overlay address plus both routed prefixes") + + require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address") + + assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey), + "allowed IPs must survive the endpoint removal unchanged") +} + +// TestRemoveEndpointAddressDoesNotScaleWithPeerCount is the regression guard for the actual +// defect: clearing one peer's endpoint used to dump and parse the whole device, so its cost +// grew with the size of the network map. On a routing peer with thousands of peers that dump +// runs on every relay and ICE transition, under the interface lock. +func TestRemoveEndpointAddressDoesNotScaleWithPeerCount(t *testing.T) { + measure := func(peerCount int) float64 { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, peerCount)[peerCount/2] + + return testing.AllocsPerRun(5, func() { + require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address") + }) + } + + small := measure(64) + large := measure(1024) + + assert.Less(t, large, small*2, + "clearing one endpoint allocated %.0f objects with 1024 peers against %.0f with 64: the cost still scales with the peer count", + large, small) +} + +// TestRemoveEndpointAddressFallsBackToDevice covers a peer the store never saw, which is what +// an out-of-band reconfiguration of the device leaves behind. The device stays the source of +// truth in that case, so the allowed IPs must still be preserved. +func TestRemoveEndpointAddressFallsBackToDevice(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 3)[1] + require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix") + + before := peerAllowedIPs(t, c, peerKey) + c.allowedIPs.reset() + + require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address") + + assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey), + "allowed IPs recovered from the device must be preserved") + + recovered, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + assert.True(t, ok, "the fallback must seed the store so the next call skips the device dump") + assert.Len(t, recovered, 2, "seeded prefixes") +} + +func TestRemoveAllowedIPKeepsTheOtherPrefixes(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 3)[0] + routed := netip.MustParsePrefix("10.20.0.0/16") + require.NoError(t, c.AddAllowedIP(peerKey, routed), "add routed prefix") + require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("192.168.7.0/24")), "add routed prefix") + + require.NoError(t, c.RemoveAllowedIP(peerKey, routed), "remove routed prefix") + + assert.ElementsMatch(t, []string{"100.64.0.0/32", "192.168.7.0/24"}, peerAllowedIPs(t, c, peerKey), + "only the removed prefix should be gone") + + assert.ErrorIs(t, c.RemoveAllowedIP(peerKey, routed), ErrAllowedIPNotFound, + "removing a prefix that is no longer configured must be reported") +} + +// TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt covers the lazy connection window documented +// in #6863: AddAllowedIP is update-only, a silent no-op when the peer is absent, so it must not +// leave the store claiming prefixes the device never took. RemoveEndpointAddress re-adds a peer +// without update-only, so a phantom entry would create a peer the device had dropped, and a +// created peer would steal those allowed IPs from whichever peer legitimately holds them. +func TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt(t *testing.T) { + c := newTestUSPConfigurer(t) + seedPeers(t, c, 2) + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + absent := priv.PublicKey().String() + + require.NoError(t, c.AddAllowedIP(absent, netip.MustParsePrefix("10.20.0.0/16")), + "update-only add on an absent peer is a silent no-op") + + stats, err := c.FullStats() + require.NoError(t, err, "read device stats") + require.Len(t, stats.Peers, 2, "the absent peer must not have been created by AddAllowedIP") + + assert.ErrorIs(t, c.RemoveEndpointAddress(absent), ErrPeerNotFound, + "clearing the endpoint of a peer the device does not have must fail") + + stats, err = c.FullStats() + require.NoError(t, err, "read device stats") + assert.Len(t, stats.Peers, 2, "no peer may be created while clearing an endpoint") +} + +// TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer covers WireGuard's rule that an +// allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away from +// whichever peer held it before. UpdatePeer relies on that rule rather than removing the prefix +// from the previous holder itself, so a prefix handed over between peers must not come back. +func TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer(t *testing.T) { + c := newTestUSPConfigurer(t) + keys := seedPeers(t, c, 2) + peerA, peerB := keys[0], keys[1] + routed := netip.MustParsePrefix("10.20.0.0/16") + + require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A") + require.Contains(t, peerAllowedIPs(t, c, peerA), routed.String(), "A must hold the prefix") + + // The route moves to B. The device takes it away from A on its own. + require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B") + require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix") + require.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), "the device must have taken it from A") + + require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint") + + assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), + "clearing A's endpoint must not take the prefix back from B") + assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), + "B must still hold the prefix") +} + +// TestPresharedKeyCreatedPeerTakesPartInPrefixHandover covers a peer created by a preshared +// key write rather than by a peer update. Rosenpass applies a peer's first key without +// updateOnly, which creates the peer on the device, so a store that ignored that operation +// would treat the peer as unknown and would not account for a prefix later handed over to it. +func TestPresharedKeyCreatedPeerTakesPartInPrefixHandover(t *testing.T) { + c := newTestUSPConfigurer(t) + peerA := seedPeers(t, c, 1)[0] + routed := netip.MustParsePrefix("10.20.0.0/16") + require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A") + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + peerB := priv.PublicKey().String() + + psk, err := wgtypes.GenerateKey() + require.NoError(t, err, "generate preshared key") + require.NoError(t, c.SetPresharedKey(peerB, psk, false), "a first key creates the peer") + + require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B") + require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix") + + require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint") + + assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), + "clearing A's endpoint must not take the prefix back from B") + assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must still hold the prefix") +} + +// TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice is the end to end form of the +// conversion: a v4-mapped prefix must not reach the device as a zero length allowed IP, +// which would route every v4 address to that peer. +func TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice(t *testing.T) { + c := newTestUSPConfigurer(t) + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + peerKey := priv.PublicKey().String() + + mapped := netip.MustParsePrefix("::ffff:10.1.2.3/112") + require.NoError(t, c.UpdatePeer(peerKey, []netip.Prefix{mapped}, 25*time.Second, nil, nil), "add peer") + + onDevice := peerAllowedIPs(t, c, peerKey) + assert.NotContains(t, onDevice, "0.0.0.0/0", "the device must not be given a catch-all allowed IP") + assert.Equal(t, []string{"10.1.0.0/16"}, onDevice, "the device holds the normalized prefix") + + recorded, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + require.True(t, ok, "the peer must be recorded") + require.Len(t, recorded, 1, "one prefix recorded") + assert.Equal(t, onDevice[0], recorded[0].String(), "device and store must agree") +} + +// TestUpdatePeerWithAnUnusableEndpointTouchesNothing pins the ordering: the endpoint is +// parsed before the device is configured, so a failure cannot leave the device holding a +// peer that the store never learned about, with the prefix handover skipped along with it. +func TestUpdatePeerWithAnUnusableEndpointTouchesNothing(t *testing.T) { + c := newTestUSPConfigurer(t) + seedPeers(t, c, 2) + + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err, "generate peer private key") + peerKey := priv.PublicKey().String() + + // A three byte address has no textual form netip can parse back. + endpoint := &net.UDPAddr{IP: net.IP{1, 2, 3}, Port: 51820} + require.Error(t, c.UpdatePeer(peerKey, []netip.Prefix{netip.MustParsePrefix("10.30.0.0/16")}, + 25*time.Second, endpoint, nil), "an unusable endpoint must fail the update") + + stats, err := c.FullStats() + require.NoError(t, err, "read device stats") + assert.Len(t, stats.Peers, 2, "the peer must not have reached the device") + + _, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + assert.False(t, ok, "the peer must not have been recorded either") +} + +// TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses covers a removal that never reached the +// device. A single peer removal is one write, so a failure leaves the peer on the device +// exactly as it was, and the record still describes it; dropping it would only force the +// next caller to read the whole device back for an answer it already had. +func TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses(t *testing.T) { + c := newTestUSPConfigurer(t) + peerKey := seedPeers(t, c, 1)[0] + require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix") + + before, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + require.True(t, ok, "the peer must be recorded before the removal") + require.Len(t, before, 2, "overlay address plus routed prefix") + + // A closed device refuses every write, which is the shape of any failed removal. + c.device.Close() + + require.Error(t, c.RemovePeer(peerKey), "the removal must report the failure") + + after, ok := c.allowedIPs.get(mustParseKey(t, peerKey)) + require.True(t, ok, "a peer still on the device must stay recorded") + assert.Equal(t, before, after, "the record must describe the peer the device kept") +} + +// mustParseKey turns the textual key the configurer API takes into the form the store +// keys on. +func mustParseKey(t *testing.T, key string) wgtypes.Key { + t.Helper() + + parsed, err := wgtypes.ParseKey(key) + require.NoError(t, err, "parse peer key") + return parsed +} diff --git a/client/iface/iface.go b/client/iface/iface.go index 247f421a2..f6006fa87 100644 --- a/client/iface/iface.go +++ b/client/iface/iface.go @@ -51,7 +51,6 @@ func ValidateMTU(mtu uint16) error { type wgProxyFactory interface { GetProxy() wgproxy.Proxy - GetProxyPort() uint16 Free() error } @@ -81,12 +80,6 @@ func (w *WGIface) GetProxy() wgproxy.Proxy { return w.wgProxyFactory.GetProxy() } -// GetProxyPort returns the proxy port used by the WireGuard proxy. -// Returns 0 if no proxy port is used (e.g., for userspace WireGuard). -func (w *WGIface) GetProxyPort() uint16 { - return w.wgProxyFactory.GetProxyPort() -} - // GetBind returns the EndpointManager userspace bind mode. func (w *WGIface) GetBind() device.EndpointManager { w.mu.Lock() diff --git a/client/iface/iface_close_test.go b/client/iface/iface_close_test.go index 171e15d0a..ea3115ec0 100644 --- a/client/iface/iface_close_test.go +++ b/client/iface/iface_close_test.go @@ -52,7 +52,6 @@ func (f *fakeTunDevice) Close() error { type fakeProxyFactory struct{} func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil } -func (fakeProxyFactory) GetProxyPort() uint16 { return 0 } func (fakeProxyFactory) Free() error { return nil } // TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock diff --git a/client/iface/iface_destroy_windows.go b/client/iface/iface_destroy_windows.go index 0bfa4e211..54c0014c4 100644 --- a/client/iface/iface_destroy_windows.go +++ b/client/iface/iface_destroy_windows.go @@ -6,27 +6,14 @@ import ( "fmt" "os/exec" - log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/internal/wincmd" ) func (w *WGIface) Destroy() error { - netshCmd := GetSystem32Command("netsh") + netshCmd := wincmd.System32("netsh") out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput() if err != nil { return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out) } return nil } - -// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it -// in the path it will return the full path of a command assuming C:\windows\system32 as the base path. -func GetSystem32Command(command string) string { - _, err := exec.LookPath(command) - if err == nil { - return command - } - - log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command) - - return "C:\\windows\\system32\\" + command + ".exe" -} diff --git a/client/iface/iface_test.go b/client/iface/iface_test.go index 89c8cd16e..cb50ca4a1 100644 --- a/client/iface/iface_test.go +++ b/client/iface/iface_test.go @@ -40,14 +40,18 @@ func init() { peerPubKey = peerPrivateKey.PublicKey().String() } +// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist +// carries for the overlay interface. These tests create their own utun device, and +// stdnet's filter probes with wgctrl every interface it is not told to skip, which +// on a userspace WireGuard platform reaches the UAPI socket of this same process. +// Declared here rather than imported because profilemanager imports this package. +var testIFaceBlackList = []string{"wt", "utun", "tun0"} + func TestWGIface_UpdateAddr(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) addr := "100.64.0.1/8" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) { func Test_CreateInterface(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1) wgIP := "10.99.99.1/32" - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, Address: wgaddr.MustParseWGAddress(wgIP), @@ -170,10 +171,7 @@ func Test_Close(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) wgIP := "10.99.99.2/32" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) wgIP := "10.99.99.2/32" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3) wgIP := "10.99.99.5/30" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, Address: wgaddr.MustParseWGAddress(wgIP), @@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) { func Test_UpdatePeer(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) wgIP := "10.99.99.9/30" - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) { func Test_RemovePeer(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) wgIP := "10.99.99.13/30" - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) { peer2wgPort := 33200 keepAlive := 1 * time.Second - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) guid := fmt.Sprintf("{%s}", uuid.New().String()) device.CustomWindowsGUIDString = strings.ToLower(guid) @@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) { guid = fmt.Sprintf("{%s}", uuid.New().String()) device.CustomWindowsGUIDString = strings.ToLower(guid) - newNet, err = stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet = stdnet.NewNet(context.Background(), testIFaceBlackList) optsPeer2 := WGIFaceOpts{ IFaceName: peer2ifaceName, @@ -568,11 +548,14 @@ func Test_ConnectPeers(t *testing.T) { if err != nil { t.Fatal(err) } - // The peers use userspace WireGuard (stdnet transport). A tight busy-loop - // here starves the wireguard-go goroutines that process the handshake, so - // poll on a ticker instead and yield the CPU between checks. WireGuard also - // only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which - // is why the overall wait can occasionally stretch to tens of seconds. + // On Linux with the kernel module both peers are kernel devices, elsewhere + // they run on wireguard-go. A tight busy-loop here would starve the + // wireguard-go goroutines that process the handshake, so poll on a ticker + // instead and yield the CPU between checks. WireGuard also only retries a + // lost handshake initiation every REKEY_TIMEOUT (5s), which is why the + // overall wait can occasionally stretch to tens of seconds. Each side sends + // its first initiation when its peer is configured, and the first one leaves + // before the other device knows the peer, so that one is always wasted. timeout := 30 * time.Second timeoutChannel := time.After(timeout) ticker := time.NewTicker(500 * time.Millisecond) @@ -590,13 +573,26 @@ func Test_ConnectPeers(t *testing.T) { select { case <-timeoutChannel: - t.Fatalf("waiting for peer handshake timeout after %s", timeout.String()) + // The counters tell whether initiations were sent at all, whether they + // arrived, and whether only one direction is working. + t.Fatalf("waiting for peer handshake timeout after %s\n%s\n%s", timeout.String(), + describePeer(peer1ifaceName, peer2Key.PublicKey().String()), + describePeer(peer2ifaceName, peer1Key.PublicKey().String())) case <-ticker.C: } } } +func describePeer(ifaceName, peerPubKey string) string { + peer, err := getPeer(ifaceName, peerPubKey) + if err != nil { + return fmt.Sprintf("%s: peer %s: %v", ifaceName, peerPubKey, err) + } + return fmt.Sprintf("%s: peer %s endpoint=%v tx=%d rx=%d last_handshake=%v", + ifaceName, peerPubKey, peer.Endpoint, peer.TransmitBytes, peer.ReceiveBytes, peer.LastHandshakeTime) +} + func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) { wg, err := wgctrl.New() if err != nil { diff --git a/client/iface/udpmux/mux.go b/client/iface/udpmux/mux.go index c5d2de4a5..68cecc953 100644 --- a/client/iface/udpmux/mux.go +++ b/client/iface/udpmux/mux.go @@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() { } if len(networks) > 0 { if m.params.Net == nil { - var err error - if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil { - m.params.Logger.Errorf("failed to get create network: %v", err) - } + m.params.Net = stdnet.NewNet(context.Background(), nil) } ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true) diff --git a/client/iface/wgproxy/ebpf/portlookup.go b/client/iface/wgproxy/ebpf/portlookup.go deleted file mode 100644 index fce8f1507..000000000 --- a/client/iface/wgproxy/ebpf/portlookup.go +++ /dev/null @@ -1,32 +0,0 @@ -package ebpf - -import ( - "fmt" - "net" -) - -var ( - portRangeStart = 3128 - portRangeEnd = portRangeStart + 100 -) - -type portLookup struct { -} - -func (pl portLookup) searchFreePort() (int, error) { - for i := portRangeStart; i <= portRangeEnd; i++ { - if pl.tryToBind(i) == nil { - return i, nil - } - } - return 0, fmt.Errorf("failed to bind free port for eBPF proxy") -} - -func (pl portLookup) tryToBind(port int) error { - l, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port)) - if err != nil { - return err - } - _ = l.Close() - return nil -} diff --git a/client/iface/wgproxy/ebpf/portlookup_test.go b/client/iface/wgproxy/ebpf/portlookup_test.go deleted file mode 100644 index a2e92fc79..000000000 --- a/client/iface/wgproxy/ebpf/portlookup_test.go +++ /dev/null @@ -1,45 +0,0 @@ -package ebpf - -import ( - "fmt" - "net" - "testing" -) - -func Test_portLookup_searchFreePort(t *testing.T) { - pl := portLookup{} - _, err := pl.searchFreePort() - if err != nil { - t.Fatal(err) - } -} - -func Test_portLookup_on_allocated(t *testing.T) { - pl := portLookup{} - - portRangeStart = 4128 - portRangeEnd = portRangeStart + 100 - - allocatedPort, err := allocatePort(portRangeStart) - if err != nil { - t.Fatal(err) - } - defer allocatedPort.Close() - - fp, err := pl.searchFreePort() - if err != nil { - t.Fatal(err) - } - - if fp != (portRangeStart + 1) { - t.Errorf("invalid free port, expected: %d, got: %d", portRangeStart+1, fp) - } -} - -func allocatePort(port int) (net.PacketConn, error) { - c, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port)) - if err != nil { - return nil, err - } - return c, err -} diff --git a/client/iface/wgproxy/ebpf/proxy.go b/client/iface/wgproxy/ebpf/proxy.go deleted file mode 100644 index 91c741c0d..000000000 --- a/client/iface/wgproxy/ebpf/proxy.go +++ /dev/null @@ -1,243 +0,0 @@ -//go:build linux && !android - -package ebpf - -import ( - "context" - "fmt" - "net" - "sync" - - "github.com/hashicorp/go-multierror" - "github.com/pion/transport/v3" - log "github.com/sirupsen/logrus" - - nberrors "github.com/netbirdio/netbird/client/errors" - "github.com/netbirdio/netbird/client/iface/bufsize" - "github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket" - "github.com/netbirdio/netbird/client/internal/ebpf" - ebpfMgr "github.com/netbirdio/netbird/client/internal/ebpf/manager" - nbnet "github.com/netbirdio/netbird/client/net" -) - -const ( - loopbackAddr = "127.0.0.1" -) - -// WGEBPFProxy definition for proxy with EBPF support -type WGEBPFProxy struct { - localWGListenPort int - proxyPort int - mtu uint16 - - ebpfManager ebpfMgr.Manager - relayedConnStore map[uint16]net.Conn - relayedConnMutex sync.Mutex - - lastUsedPort uint16 - rawConnIPv4 net.PacketConn - rawConnIPv6 net.PacketConn - conn transport.UDPConn - - ctx context.Context - ctxCancel context.CancelFunc -} - -// NewWGEBPFProxy create new WGEBPFProxy instance -func NewWGEBPFProxy(wgPort int, mtu uint16) *WGEBPFProxy { - log.Debugf("instantiate ebpf proxy") - wgProxy := &WGEBPFProxy{ - localWGListenPort: wgPort, - mtu: mtu, - ebpfManager: ebpf.GetEbpfManagerInstance(), - relayedConnStore: make(map[uint16]net.Conn), - } - return wgProxy -} - -// Listen load ebpf program and listen the proxy -func (p *WGEBPFProxy) Listen() error { - pl := portLookup{} - proxyPort, err := pl.searchFreePort() - if err != nil { - return err - } - p.proxyPort = proxyPort - - // Prepare IPv4 raw socket (required) - p.rawConnIPv4, err = rawsocket.PrepareSenderRawSocketIPv4() - if err != nil { - return err - } - - // Prepare IPv6 raw socket (optional) - p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6() - if err != nil { - log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err) - } - - err = p.ebpfManager.LoadWgProxy(proxyPort, p.localWGListenPort) - if err != nil { - if closeErr := p.rawConnIPv4.Close(); closeErr != nil { - log.Warnf("failed to close IPv4 raw socket: %v", closeErr) - } - if p.rawConnIPv6 != nil { - if closeErr := p.rawConnIPv6.Close(); closeErr != nil { - log.Warnf("failed to close IPv6 raw socket: %v", closeErr) - } - } - return err - } - - addr := net.UDPAddr{ - Port: proxyPort, - IP: net.ParseIP(loopbackAddr), - } - - p.ctx, p.ctxCancel = context.WithCancel(context.Background()) - - conn, err := nbnet.ListenUDP("udp", &addr) - if err != nil { - if cErr := p.Free(); cErr != nil { - log.Errorf("Failed to close the wgproxy: %s", cErr) - } - return err - } - p.conn = conn - - go p.proxyToRemote() - log.Infof("local wg proxy listening on: %d", proxyPort) - return nil -} - -// AddRelayedConn add new relayed connection for the proxy -func (p *WGEBPFProxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, error) { - wgEndpointPort, err := p.storeRelayedConn(relayedConn) - if err != nil { - return nil, err - } - - log.Infof("relayed conn added to wg proxy store: %s, endpoint port: :%d", relayedConn.RemoteAddr(), wgEndpointPort) - - wgEndpoint := &net.UDPAddr{ - IP: net.ParseIP(loopbackAddr), - Port: int(wgEndpointPort), - } - return wgEndpoint, nil -} - -// Free resources except the remoteConns will be keep open. -func (p *WGEBPFProxy) Free() error { - log.Debugf("free up ebpf wg proxy") - if p.ctx != nil && p.ctx.Err() != nil { - //nolint - return nil - } - - p.ctxCancel() - - var result *multierror.Error - if p.conn != nil { - if err := p.conn.Close(); err != nil { - result = multierror.Append(result, err) - } - } - - if err := p.ebpfManager.FreeWGProxy(); err != nil { - result = multierror.Append(result, err) - } - - if p.rawConnIPv4 != nil { - if err := p.rawConnIPv4.Close(); err != nil { - result = multierror.Append(result, err) - } - } - - if p.rawConnIPv6 != nil { - if err := p.rawConnIPv6.Close(); err != nil { - result = multierror.Append(result, err) - } - } - return nberrors.FormatErrorOrNil(result) -} - -// GetProxyPort returns the proxy listening port. -func (p *WGEBPFProxy) GetProxyPort() uint16 { - return uint16(p.proxyPort) -} - -// proxyToRemote read messages from local WireGuard interface and forward it to remote conn -// From this go routine has only one instance. -func (p *WGEBPFProxy) proxyToRemote() { - buf := make([]byte, p.mtu+bufsize.WGBufferOverhead) - for p.ctx.Err() == nil { - if err := p.readAndForwardPacket(buf); err != nil { - if p.ctx.Err() != nil { - return - } - log.Errorf("failed to proxy packet to remote conn: %s", err) - } - } -} - -func (p *WGEBPFProxy) readAndForwardPacket(buf []byte) error { - n, addr, err := p.conn.ReadFromUDP(buf) - if err != nil { - return fmt.Errorf("failed to read UDP packet from WG: %w", err) - } - - p.relayedConnMutex.Lock() - conn, ok := p.relayedConnStore[uint16(addr.Port)] - p.relayedConnMutex.Unlock() - if !ok { - if p.ctx.Err() == nil { - log.Debugf("relayed conn not found by port because conn already has been closed: %d", addr.Port) - } - return nil - } - - if _, err := conn.Write(buf[:n]); err != nil { - return fmt.Errorf("forward local WG packet (%d) to remote relayed conn: %w", addr.Port, err) - } - return nil -} - -func (p *WGEBPFProxy) storeRelayedConn(relayedConn net.Conn) (uint16, error) { - p.relayedConnMutex.Lock() - defer p.relayedConnMutex.Unlock() - - np, err := p.nextFreePort() - if err != nil { - return np, err - } - p.relayedConnStore[np] = relayedConn - return np, nil -} - -func (p *WGEBPFProxy) removeRelayedConn(relayedConnID uint16) { - p.relayedConnMutex.Lock() - defer p.relayedConnMutex.Unlock() - - _, ok := p.relayedConnStore[relayedConnID] - if ok { - log.Debugf("remove relayed conn from store by port: %d", relayedConnID) - } - delete(p.relayedConnStore, relayedConnID) -} - -func (p *WGEBPFProxy) nextFreePort() (uint16, error) { - if len(p.relayedConnStore) == 65535 { - return 0, fmt.Errorf("reached maximum relayed connection numbers") - } -generatePort: - if p.lastUsedPort == 65535 { - p.lastUsedPort = 1 - } else { - p.lastUsedPort++ - } - - if _, ok := p.relayedConnStore[p.lastUsedPort]; ok { - goto generatePort - } - return p.lastUsedPort, nil -} diff --git a/client/iface/wgproxy/ebpf/proxy_test.go b/client/iface/wgproxy/ebpf/proxy_test.go deleted file mode 100644 index 228c06c9b..000000000 --- a/client/iface/wgproxy/ebpf/proxy_test.go +++ /dev/null @@ -1,56 +0,0 @@ -//go:build linux && !android - -package ebpf - -import ( - "testing" -) - -func TestWGEBPFProxy_connStore(t *testing.T) { - wgProxy := NewWGEBPFProxy(1, 1280) - - p, _ := wgProxy.storeRelayedConn(nil) - if p != 1 { - t.Errorf("invalid initial port: %d", wgProxy.lastUsedPort) - } - - numOfConns := 10 - for i := 0; i < numOfConns; i++ { - p, _ = wgProxy.storeRelayedConn(nil) - } - if p != uint16(numOfConns)+1 { - t.Errorf("invalid last used port: %d, expected: %d", p, numOfConns+1) - } - if len(wgProxy.relayedConnStore) != numOfConns+1 { - t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), numOfConns+1) - } -} - -func TestWGEBPFProxy_portCalculation_overflow(t *testing.T) { - wgProxy := NewWGEBPFProxy(1, 1280) - - _, _ = wgProxy.storeRelayedConn(nil) - wgProxy.lastUsedPort = 65535 - p, _ := wgProxy.storeRelayedConn(nil) - - if len(wgProxy.relayedConnStore) != 2 { - t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), 2) - } - - if p != 2 { - t.Errorf("invalid last used port: %d, expected: %d", p, 2) - } -} - -func TestWGEBPFProxy_portCalculation_maxConn(t *testing.T) { - wgProxy := NewWGEBPFProxy(1, 1280) - - for i := 0; i < 65535; i++ { - _, _ = wgProxy.storeRelayedConn(nil) - } - - _, err := wgProxy.storeRelayedConn(nil) - if err == nil { - t.Errorf("invalid relayed conn store calculation") - } -} diff --git a/client/iface/wgproxy/factory_kernel.go b/client/iface/wgproxy/factory_kernel.go index 7821df3de..0b2329b96 100644 --- a/client/iface/wgproxy/factory_kernel.go +++ b/client/iface/wgproxy/factory_kernel.go @@ -8,11 +8,13 @@ import ( log "github.com/sirupsen/logrus" - "github.com/netbirdio/netbird/client/iface/wgproxy/ebpf" + "github.com/netbirdio/netbird/client/iface/wgproxy/loopback" udpProxy "github.com/netbirdio/netbird/client/iface/wgproxy/udp" ) const ( + envDisableKernelWGProxy = "NB_DISABLE_KERNEL_WG_PROXY" + // envDisableEBPFWGProxy is a deprecated alias for envDisableKernelWGProxy. envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY" ) @@ -20,7 +22,7 @@ type KernelFactory struct { wgPort int mtu uint16 - ebpfProxy *ebpf.WGEBPFProxy + loopbackProxy *loopback.Proxy } func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory { @@ -29,55 +31,56 @@ func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory { mtu: mtu, } - if isEBPFDisabled() { + if isKernelProxyDisabled() { log.Infof("WireGuard Proxy Factory will produce UDP proxy") - log.Infof("eBPF WireGuard proxy is disabled via %s environment variable", envDisableEBPFWGProxy) return f } - ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, mtu) - if err := ebpfProxy.Listen(); err != nil { + loopbackProxy := loopback.NewProxy(wgPort, mtu) + if err := loopbackProxy.Listen(); err != nil { log.Infof("WireGuard Proxy Factory will produce UDP proxy") - log.Warnf("failed to initialize ebpf proxy, fallback to user space proxy: %s", err) + log.Warnf("failed to initialize loopback proxy, fallback to user space proxy: %s", err) return f } - log.Infof("WireGuard Proxy Factory will produce eBPF proxy") - f.ebpfProxy = ebpfProxy + log.Infof("WireGuard Proxy Factory will produce loopback proxy") + f.loopbackProxy = loopbackProxy return f } func (w *KernelFactory) GetProxy() Proxy { - if w.ebpfProxy == nil { + if w.loopbackProxy == nil { return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu) } - return ebpf.NewProxyWrapper(w.ebpfProxy) -} - -// GetProxyPort returns the eBPF proxy port, or 0 if eBPF is not active. -func (w *KernelFactory) GetProxyPort() uint16 { - if w.ebpfProxy == nil { - return 0 - } - return w.ebpfProxy.GetProxyPort() + return loopback.NewProxyWrapper(w.loopbackProxy) } func (w *KernelFactory) Free() error { - if w.ebpfProxy == nil { + if w.loopbackProxy == nil { return nil } - return w.ebpfProxy.Free() + return w.loopbackProxy.Free() } -func isEBPFDisabled() bool { - val := os.Getenv(envDisableEBPFWGProxy) +func isKernelProxyDisabled() bool { + env := envDisableKernelWGProxy + val := os.Getenv(env) + if val == "" { + env = envDisableEBPFWGProxy + val = os.Getenv(env) + } if val == "" { return false } + disabled, err := strconv.ParseBool(val) if err != nil { - log.Warnf("failed to parse %s: %v", envDisableEBPFWGProxy, err) + log.Warnf("failed to parse %s: %v", env, err) return false } + + if disabled { + log.Infof("kernel WireGuard proxy is disabled via %s", env) + } return disabled } diff --git a/client/iface/wgproxy/factory_usp.go b/client/iface/wgproxy/factory_usp.go index bbd67e076..a1b1c34d7 100644 --- a/client/iface/wgproxy/factory_usp.go +++ b/client/iface/wgproxy/factory_usp.go @@ -24,11 +24,6 @@ func (w *USPFactory) GetProxy() Proxy { return proxyBind.NewProxyBind(w.bind, w.mtu) } -// GetProxyPort returns 0 as userspace WireGuard doesn't use a separate proxy port. -func (w *USPFactory) GetProxyPort() uint16 { - return 0 -} - func (w *USPFactory) Free() error { return nil } diff --git a/client/iface/wgproxy/loopback/addr.go b/client/iface/wgproxy/loopback/addr.go new file mode 100644 index 000000000..52feee295 --- /dev/null +++ b/client/iface/wgproxy/loopback/addr.go @@ -0,0 +1,70 @@ +//go:build linux && !android + +package loopback + +import ( + "fmt" + "net/netip" +) + +// Peer endpoints live in the upper half of 127.0.0.0/8. Everything in that +// range is delivered to the loopback device without any address or route being +// configured, and staying out of 127.0.0.0/9 keeps well-known squatters such as +// 127.0.0.53 (systemd-resolved) and 127.0.1.1 out of the way. +const ( + addrRangeBase uint32 = 0x7f800000 // 127.128.0.0 + addrRangeSize uint32 = 1 << 23 // /9 + addrRangePrefix = "127.128.0.0/9" +) + +// allocator hands out one loopback address per relayed connection. The address +// is the peer's identity: WireGuard sends to it, and the proxy recovers which +// peer a packet belongs to from the destination address. +type allocator struct { + cursor uint32 +} + +// next returns the first free address at or after the cursor, wrapping once. +// inUse reports whether an address is already handed out. +func (a *allocator) next(inUse func(netip.Addr) bool) (netip.Addr, error) { + for i := uint32(0); i < addrRangeSize; i++ { + a.cursor = (a.cursor + 1) % addrRangeSize + addr := addrFromOffset(a.cursor) + if !addr.IsValid() { + continue + } + if inUse(addr) { + continue + } + return addr, nil + } + return netip.Addr{}, fmt.Errorf("no free endpoint address in %s", addrRangePrefix) +} + +// addrFromOffset maps an offset in the range to an address, skipping the .0 and +// .255 hosts. They are unremarkable on loopback, but tools and firewall rules +// tend to treat them as network and broadcast addresses. +func addrFromOffset(offset uint32) netip.Addr { + last := offset & 0xff + if last == 0 || last == 0xff { + return netip.Addr{} + } + + v := addrRangeBase + offset + return netip.AddrFrom4([4]byte{ + byte(v >> 24), + byte(v >> 16), + byte(v >> 8), + byte(v), + }) +} + +// inRange reports whether addr is one this proxy could have handed out. +func inRange(addr netip.Addr) bool { + if !addr.Is4() { + return false + } + b := addr.As4() + v := uint32(b[0])<<24 | uint32(b[1])<<16 | uint32(b[2])<<8 | uint32(b[3]) + return v >= addrRangeBase && v < addrRangeBase+addrRangeSize && b[3] != 0 && b[3] != 0xff +} diff --git a/client/iface/wgproxy/loopback/addr_test.go b/client/iface/wgproxy/loopback/addr_test.go new file mode 100644 index 000000000..3755b7269 --- /dev/null +++ b/client/iface/wgproxy/loopback/addr_test.go @@ -0,0 +1,114 @@ +//go:build linux && !android + +package loopback + +import ( + "net/netip" + "testing" +) + +func TestAllocatorHandsOutDistinctAddresses(t *testing.T) { + var a allocator + taken := make(map[netip.Addr]bool) + + for i := 0; i < 1000; i++ { + addr, err := a.next(func(candidate netip.Addr) bool { return taken[candidate] }) + if err != nil { + t.Fatalf("allocate %d: %v", i, err) + } + if taken[addr] { + t.Fatalf("address %s handed out twice", addr) + } + if !inRange(addr) { + t.Fatalf("address %s outside %s", addr, addrRangePrefix) + } + taken[addr] = true + } +} + +func TestAllocatorSkipsNetworkAndBroadcastHosts(t *testing.T) { + var a allocator + taken := make(map[netip.Addr]bool) + + // enough allocations to walk past a .255/.0 boundary + for i := 0; i < 600; i++ { + addr, err := a.next(func(candidate netip.Addr) bool { return taken[candidate] }) + if err != nil { + t.Fatalf("allocate %d: %v", i, err) + } + last := addr.As4()[3] + if last == 0 || last == 255 { + t.Fatalf("address %s ends in .%d", addr, last) + } + taken[addr] = true + } +} + +func TestAllocatorReusesReleasedAddresses(t *testing.T) { + var a allocator + taken := make(map[netip.Addr]bool) + inUse := func(candidate netip.Addr) bool { return taken[candidate] } + alloc := func() netip.Addr { + t.Helper() + addr, err := a.next(inUse) + if err != nil { + t.Fatalf("allocate: %v", err) + } + taken[addr] = true + return addr + } + + first := alloc() + second := alloc() + delete(taken, first) + + // The cursor only moves forward, so a released address comes back after a + // wrap. Park the cursor near the end of the range instead of allocating + // 2^23 addresses: the next call takes the last usable address, and the one + // after that wraps past the skipped .255 and .0 hosts to the released one. + a.cursor = addrRangeSize - 3 + last := alloc() + if want := netip.MustParseAddr("127.255.255.254"); last != want { + t.Fatalf("expected the last usable address %s before the wrap, got %s", want, last) + } + + if reused := alloc(); reused != first { + t.Fatalf("expected the released address %s after the wrap, got %s", first, reused) + } + + // second is still held, so the allocator must step over it. + if next := alloc(); next == second { + t.Fatalf("allocator handed out %s while it was still in use", second) + } +} + +func TestInRange(t *testing.T) { + tests := []struct { + addr string + want bool + }{ + {"127.128.0.1", true}, + {"127.255.255.254", true}, + {"127.128.0.0", false}, // network host, never handed out + {"127.128.5.255", false}, // broadcast host, never handed out + {"127.127.255.255", false}, // below the range, where 127.0.0.53 and friends live + {"127.0.0.1", false}, + {"127.0.0.53", false}, + {"127.0.1.1", false}, + {"128.0.0.1", false}, + {"10.0.0.1", false}, + } + + for _, tc := range tests { + addr := netip.MustParseAddr(tc.addr) + if got := inRange(addr); got != tc.want { + t.Errorf("inRange(%s) = %v, want %v", tc.addr, got, tc.want) + } + } +} + +func TestInRangeIgnoresIPv6(t *testing.T) { + if inRange(netip.MustParseAddr("::1")) { + t.Error("inRange(::1) = true, want false") + } +} diff --git a/client/iface/wgproxy/loopback/proxy.go b/client/iface/wgproxy/loopback/proxy.go new file mode 100644 index 000000000..f9364c766 --- /dev/null +++ b/client/iface/wgproxy/loopback/proxy.go @@ -0,0 +1,291 @@ +//go:build linux && !android + +package loopback + +import ( + "context" + "fmt" + "net" + "net/netip" + "sync" + "syscall" + + "github.com/hashicorp/go-multierror" + log "github.com/sirupsen/logrus" + "golang.org/x/net/ipv4" + "golang.org/x/sys/unix" + + nberrors "github.com/netbirdio/netbird/client/errors" + "github.com/netbirdio/netbird/client/iface/bufsize" + "github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket" +) + +const ( + loopbackDevice = "lo" + + portRangeStart = 3128 + portRangeEnd = portRangeStart + 100 +) + +// Proxy forwards packets between relayed connections and a local kernel +// WireGuard instance. Every relayed peer gets its own loopback address as its +// WireGuard endpoint, so a single socket serves all of them: the destination +// address of an incoming packet identifies the peer. +type Proxy struct { + localWGListenPort int + mtu uint16 + proxyPort int + + conn *net.UDPConn + packetConn *ipv4.PacketConn + loIndex int + rawConnIPv4 net.PacketConn + rawConnIPv6 net.PacketConn + + relayedConnMutex sync.Mutex + relayedConnStore map[netip.Addr]net.Conn + addrs allocator + + ctx context.Context + ctxCancel context.CancelFunc +} + +// NewProxy creates a proxy for the WireGuard instance listening on wgPort. +func NewProxy(wgPort int, mtu uint16) *Proxy { + log.Debugf("instantiate loopback wg proxy") + return &Proxy{ + localWGListenPort: wgPort, + mtu: mtu, + relayedConnStore: make(map[netip.Addr]net.Conn), + } +} + +// Listen opens the shared socket and starts forwarding WireGuard packets to the +// relayed connections. +func (p *Proxy) Listen() error { + rawConnIPv4, err := rawsocket.PrepareSenderRawSocketIPv4() + if err != nil { + return fmt.Errorf("prepare IPv4 raw socket: %w", err) + } + p.rawConnIPv4 = rawConnIPv4 + + p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6() + if err != nil { + log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err) + } + + loopback, err := net.InterfaceByName(loopbackDevice) + if err != nil { + if freeErr := p.Free(); freeErr != nil { + log.Errorf("failed to free the wgproxy: %s", freeErr) + } + return fmt.Errorf("look up %s: %w", loopbackDevice, err) + } + p.loIndex = loopback.Index + + if err := p.listen(); err != nil { + if freeErr := p.Free(); freeErr != nil { + log.Errorf("failed to free the wgproxy: %s", freeErr) + } + return err + } + + p.ctx, p.ctxCancel = context.WithCancel(context.Background()) + + go p.proxyToRemote() + log.Infof("local wg proxy listening on %s:%d", addrRangePrefix, p.proxyPort) + return nil +} + +// listen binds the shared socket on the first free port of the range. The bind +// has to be a wildcard one to receive every peer address in the range, so it is +// restricted to the loopback device: without that the port would be reachable +// on every interface. +func (p *Proxy) listen() error { + var lastErr error + for port := portRangeStart; port <= portRangeEnd; port++ { + err := p.listenOn(port) + if err == nil { + p.proxyPort = port + return nil + } + lastErr = err + } + return fmt.Errorf("bind proxy port in range %d-%d: %w", portRangeStart, portRangeEnd, lastErr) +} + +func (p *Proxy) listenOn(proxyPort int) error { + lc := net.ListenConfig{ + Control: func(_, _ string, c syscall.RawConn) error { + var sockErr error + if err := c.Control(func(fd uintptr) { + if err := unix.SetsockoptString(int(fd), unix.SOL_SOCKET, unix.SO_BINDTODEVICE, loopbackDevice); err != nil { + sockErr = fmt.Errorf("bind to %s: %w", loopbackDevice, err) + return + } + }); err != nil { + return fmt.Errorf("control socket: %w", err) + } + return sockErr + }, + } + + conn, err := lc.ListenPacket(context.Background(), "udp4", fmt.Sprintf(":%d", proxyPort)) + if err != nil { + return fmt.Errorf("listen on :%d: %w", proxyPort, err) + } + + udpConn, ok := conn.(*net.UDPConn) + if !ok { + if closeErr := conn.Close(); closeErr != nil { + log.Errorf("failed to close proxy conn: %s", closeErr) + } + return fmt.Errorf("unexpected conn type %T", conn) + } + + packetConn := ipv4.NewPacketConn(udpConn) + // the destination address carries the peer identity, the interface index is + // checked on receive as a second line of defense behind SO_BINDTODEVICE + if err := packetConn.SetControlMessage(ipv4.FlagDst|ipv4.FlagInterface, true); err != nil { + if closeErr := udpConn.Close(); closeErr != nil { + log.Errorf("failed to close proxy conn: %s", closeErr) + } + return fmt.Errorf("request destination address: %w", err) + } + + p.conn = udpConn + p.packetConn = packetConn + return nil +} + +// AddRelayedConn assigns an endpoint address to the relayed connection and +// returns the address WireGuard should send to, along with the key the +// connection is stored under. +func (p *Proxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, netip.Addr, error) { + addr, err := p.storeRelayedConn(relayedConn) + if err != nil { + return nil, netip.Addr{}, err + } + + log.Infof("relayed conn added to wg proxy store: %s, endpoint address: %s", relayedConn.RemoteAddr(), addr) + + return &net.UDPAddr{ + IP: addr.AsSlice(), + Port: p.proxyPort, + }, addr, nil +} + +// Free releases the proxy resources. The relayed connections are left open. +func (p *Proxy) Free() error { + log.Debugf("free up loopback wg proxy") + if p.ctx != nil && p.ctx.Err() != nil { + //nolint + return nil + } + + if p.ctxCancel != nil { + p.ctxCancel() + } + + var result *multierror.Error + if p.conn != nil { + if err := p.conn.Close(); err != nil { + result = multierror.Append(result, err) + } + } + + if p.rawConnIPv4 != nil { + if err := p.rawConnIPv4.Close(); err != nil { + result = multierror.Append(result, err) + } + } + + if p.rawConnIPv6 != nil { + if err := p.rawConnIPv6.Close(); err != nil { + result = multierror.Append(result, err) + } + } + return nberrors.FormatErrorOrNil(result) +} + +// proxyToRemote reads packets from the local WireGuard instance and forwards +// them to the relayed connection the destination address belongs to. +func (p *Proxy) proxyToRemote() { + buf := make([]byte, p.mtu+bufsize.WGBufferOverhead) + for p.ctx.Err() == nil { + if err := p.readAndForwardPacket(buf); err != nil { + if p.ctx.Err() != nil { + return + } + log.Errorf("failed to proxy packet to remote conn: %s", err) + } + } +} + +func (p *Proxy) readAndForwardPacket(buf []byte) error { + n, cm, _, err := p.packetConn.ReadFrom(buf) + if err != nil { + return fmt.Errorf("read UDP packet from WG: %w", err) + } + + if cm == nil { + return fmt.Errorf("no control message on packet") + } + + if cm.IfIndex != p.loIndex { + log.Tracef("dropping packet received on interface %d instead of %s", cm.IfIndex, loopbackDevice) + return nil + } + + dst, ok := netip.AddrFromSlice(cm.Dst.To4()) + if !ok || !inRange(dst) { + log.Tracef("dropping packet for unexpected destination %s", cm.Dst) + return nil + } + + p.relayedConnMutex.Lock() + conn, ok := p.relayedConnStore[dst] + p.relayedConnMutex.Unlock() + if !ok { + if p.ctx.Err() == nil { + log.Debugf("relayed conn not found by address because conn already has been closed: %s", dst) + } + return nil + } + + if _, err := conn.Write(buf[:n]); err != nil { + return fmt.Errorf("forward local WG packet (%s) to remote relayed conn: %w", dst, err) + } + return nil +} + +func (p *Proxy) storeRelayedConn(relayedConn net.Conn) (netip.Addr, error) { + p.relayedConnMutex.Lock() + defer p.relayedConnMutex.Unlock() + + addr, err := p.addrs.next(func(a netip.Addr) bool { + _, ok := p.relayedConnStore[a] + return ok + }) + if err != nil { + return netip.Addr{}, err + } + + p.relayedConnStore[addr] = relayedConn + return addr, nil +} + +// removeRelayedConn releases an endpoint address. It only removes the entry +// while it still belongs to relayedConn, so a late release cannot take an +// address away from the peer it was handed to next. +func (p *Proxy) removeRelayedConn(addr netip.Addr, relayedConn net.Conn) { + p.relayedConnMutex.Lock() + defer p.relayedConnMutex.Unlock() + + if stored, ok := p.relayedConnStore[addr]; !ok || stored != relayedConn { + return + } + + log.Debugf("remove relayed conn from store by address: %s", addr) + delete(p.relayedConnStore, addr) +} diff --git a/client/iface/wgproxy/loopback/proxy_privileged_test.go b/client/iface/wgproxy/loopback/proxy_privileged_test.go new file mode 100644 index 000000000..6314fe2de --- /dev/null +++ b/client/iface/wgproxy/loopback/proxy_privileged_test.go @@ -0,0 +1,196 @@ +//go:build linux && !android && privileged + +package loopback + +import ( + "context" + "net" + "strconv" + "testing" + "time" +) + +const testWGPort = 51862 + +// relayEnd stands in for a relayed connection: the proxy writes what it read +// from WireGuard into it, and the test reads it back out here. +func relayEnd(t *testing.T) (proxySide net.Conn, testSide *net.UDPConn) { + t.Helper() + + testSide, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatalf("relay listener: %v", err) + } + t.Cleanup(func() { + if err := testSide.Close(); err != nil { + t.Logf("close relay listener: %v", err) + } + }) + + proxySide, err = net.Dial("udp", testSide.LocalAddr().String()) + if err != nil { + t.Fatalf("relay conn: %v", err) + } + t.Cleanup(func() { + if err := proxySide.Close(); err != nil { + t.Logf("close relay conn: %v", err) + } + }) + + return proxySide, testSide +} + +// TestProxyDemuxesByDestinationAddress is the core of the design: one socket +// serves every peer, and the destination address decides which relayed +// connection a WireGuard packet belongs to. +func TestProxyDemuxesByDestinationAddress(t *testing.T) { + proxy := NewProxy(testWGPort, 1280) + if err := proxy.Listen(); err != nil { + t.Fatalf("listen: %v", err) + } + defer func() { + if err := proxy.Free(); err != nil { + t.Errorf("free proxy: %v", err) + } + }() + + const peers = 3 + endpoints := make([]*net.UDPAddr, 0, peers) + readers := make([]*net.UDPConn, 0, peers) + for i := 0; i < peers; i++ { + proxySide, testSide := relayEnd(t) + endpoint, _, err := proxy.AddRelayedConn(proxySide) + if err != nil { + t.Fatalf("add relayed conn %d: %v", i, err) + } + if endpoint.Port != proxy.proxyPort { + t.Errorf("peer %d endpoint port = %d, want the shared proxy port %d", i, endpoint.Port, proxy.proxyPort) + } + endpoints = append(endpoints, endpoint) + readers = append(readers, testSide) + } + + // every peer must have its own address, otherwise they are indistinguishable + seen := make(map[string]bool, peers) + for i, endpoint := range endpoints { + if seen[endpoint.IP.String()] { + t.Fatalf("peer %d reuses endpoint address %s", i, endpoint.IP) + } + seen[endpoint.IP.String()] = true + } + + wgSock, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatalf("wg socket: %v", err) + } + defer func() { + if err := wgSock.Close(); err != nil { + t.Logf("close wg socket: %v", err) + } + }() + + for i, endpoint := range endpoints { + payload := []byte{byte(i), 'p', 'k', 't'} + if _, err := wgSock.WriteTo(payload, endpoint); err != nil { + t.Fatalf("write to peer %d endpoint %s: %v", i, endpoint, err) + } + + buf := make([]byte, 1500) + if err := readers[i].SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + n, _, err := readers[i].ReadFrom(buf) + if err != nil { + t.Fatalf("peer %d did not receive its packet: %v", i, err) + } + if string(buf[:n]) != string(payload) { + t.Errorf("peer %d got %q, want %q", i, buf[:n], payload) + } + + // no other peer may see it + for j, other := range readers { + if j == i { + continue + } + if err := other.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + if _, _, err := other.ReadFrom(buf); err == nil { + t.Errorf("packet for peer %d also delivered to peer %d", i, j) + } + } + } +} + +// TestProxyDropsPacketsOutsideTheRange guards the wildcard bind: anything that +// is not addressed to a handed-out endpoint must not reach a relayed peer. +func TestProxyDropsPacketsOutsideTheRange(t *testing.T) { + proxy := NewProxy(testWGPort+1, 1280) + if err := proxy.Listen(); err != nil { + t.Fatalf("listen: %v", err) + } + defer func() { + if err := proxy.Free(); err != nil { + t.Errorf("free proxy: %v", err) + } + }() + + proxySide, testSide := relayEnd(t) + if _, _, err := proxy.AddRelayedConn(proxySide); err != nil { + t.Fatalf("add relayed conn: %v", err) + } + + sender, err := net.Dial("udp", net.JoinHostPort("127.0.0.1", strconv.Itoa(proxy.proxyPort))) + if err != nil { + t.Fatalf("sender: %v", err) + } + defer func() { + if err := sender.Close(); err != nil { + t.Logf("close sender: %v", err) + } + }() + + if _, err := sender.Write([]byte("stray")); err != nil { + t.Fatalf("write stray packet: %v", err) + } + + buf := make([]byte, 1500) + if err := testSide.SetReadDeadline(time.Now().Add(500 * time.Millisecond)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + if _, _, err := testSide.ReadFrom(buf); err == nil { + t.Error("packet addressed to 127.0.0.1 was forwarded to a relayed peer") + } +} + +// A wrapper that is closed before it starts forwarding still has to give its +// endpoint address back, otherwise the range leaks an address per attempt. +func TestClosingBeforeWorkReleasesTheAddress(t *testing.T) { + proxy := NewProxy(testWGPort+2, 1280) + if err := proxy.Listen(); err != nil { + t.Fatalf("listen: %v", err) + } + defer func() { + if err := proxy.Free(); err != nil { + t.Errorf("free proxy: %v", err) + } + }() + + proxySide, _ := relayEnd(t) + wrapper := NewProxyWrapper(proxy) + if err := wrapper.AddRelayedConn(context.Background(), nil, proxySide); err != nil { + t.Fatalf("add relayed conn: %v", err) + } + + if got := len(proxy.relayedConnStore); got != 1 { + t.Fatalf("store holds %d entries after adding one conn, want 1", got) + } + + if err := wrapper.CloseConn(); err != nil { + t.Fatalf("close conn: %v", err) + } + + if got := len(proxy.relayedConnStore); got != 0 { + t.Errorf("store holds %d entries after close, want 0", got) + } +} diff --git a/client/iface/wgproxy/ebpf/wrapper.go b/client/iface/wgproxy/loopback/wrapper.go similarity index 84% rename from client/iface/wgproxy/ebpf/wrapper.go rename to client/iface/wgproxy/loopback/wrapper.go index f75e21aa6..a9cb1ab59 100644 --- a/client/iface/wgproxy/ebpf/wrapper.go +++ b/client/iface/wgproxy/loopback/wrapper.go @@ -1,6 +1,6 @@ //go:build linux && !android -package ebpf +package loopback import ( "context" @@ -8,6 +8,7 @@ import ( "fmt" "io" "net" + "net/netip" "sync" "github.com/google/gopacket" @@ -95,13 +96,14 @@ func NewPacketHeaders(localWGListenPort int, endpoint *net.UDPAddr) (*PacketHead // ProxyWrapper help to keep the remoteConn instance for net.Conn.Close function call type ProxyWrapper struct { - wgeBPFProxy *WGEBPFProxy + proxy *Proxy remoteConn net.Conn ctx context.Context cancel context.CancelFunc wgRelayedEndpointAddr *net.UDPAddr + peerAddr netip.Addr headers *PacketHeaders headerCurrentUsed *PacketHeaders rawConn net.PacketConn @@ -113,36 +115,44 @@ type ProxyWrapper struct { closeListener *listener.CloseListener } -func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper { +func NewProxyWrapper(proxy *Proxy) *ProxyWrapper { return &ProxyWrapper{ - wgeBPFProxy: proxy, + proxy: proxy, pausedCond: sync.NewCond(&sync.Mutex{}), closeListener: listener.NewCloseListener(), } } func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { - addr, err := p.wgeBPFProxy.AddRelayedConn(remoteConn) + addr, peerAddr, err := p.proxy.AddRelayedConn(remoteConn) if err != nil { return fmt.Errorf("add relayed conn: %w", err) } - headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr) + // the endpoint address is otherwise only released by the forwarding + // goroutine, which never starts when the setup below fails + release := func() { p.proxy.removeRelayedConn(peerAddr, remoteConn) } + + headers, err := NewPacketHeaders(p.proxy.localWGListenPort, addr) if err != nil { + release() return fmt.Errorf("create packet sender: %w", err) } // Check if required raw connection is available - if !headers.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil { + if !headers.isIPv4 && p.proxy.rawConnIPv6 == nil { + release() return errIPv6ConnNotAvailable } - if headers.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil { + if headers.isIPv4 && p.proxy.rawConnIPv4 == nil { + release() return errIPv4ConnNotAvailable } p.remoteConn = remoteConn p.ctx, p.cancel = context.WithCancel(ctx) p.wgRelayedEndpointAddr = addr + p.peerAddr = peerAddr p.headers = headers p.rawConn = p.selectRawConn(headers) return nil @@ -193,18 +203,18 @@ func (p *ProxyWrapper) RedirectAs(endpoint *net.UDPAddr) { return } - header, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, endpoint) + header, err := NewPacketHeaders(p.proxy.localWGListenPort, endpoint) if err != nil { log.Errorf("failed to create packet headers: %s", err) return } // Check if required raw connection is available - if !header.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil { + if !header.isIPv4 && p.proxy.rawConnIPv6 == nil { log.Error(errIPv6ConnNotAvailable) return } - if header.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil { + if header.isIPv4 && p.proxy.rawConnIPv4 == nil { log.Error(errIPv4ConnNotAvailable) return } @@ -240,6 +250,10 @@ func (p *ProxyWrapper) CloseConn() error { p.closeListener.SetCloseListener(nil) + // releases the endpoint address for a wrapper that was never started, and + // is a no-op once the forwarding goroutine has released it + p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn) + p.pausedCond.L.Lock() p.paused = false p.pausedCond.Signal() @@ -252,9 +266,9 @@ func (p *ProxyWrapper) CloseConn() error { } func (p *ProxyWrapper) proxyToLocal(ctx context.Context) { - defer p.wgeBPFProxy.removeRelayedConn(uint16(p.wgRelayedEndpointAddr.Port)) + defer p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn) - buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead) + buf := make([]byte, p.proxy.mtu+bufsize.WGBufferOverhead) for { n, err := p.readFromRemote(ctx, buf) if err != nil { @@ -286,7 +300,7 @@ func (p *ProxyWrapper) readFromRemote(ctx context.Context, buf []byte) (int, err } p.closeListener.Notify() if !errors.Is(err, io.EOF) { - log.Errorf("failed to read from relayed conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err) + log.Errorf("failed to read from relayed conn (endpoint: %s): %s", p.wgRelayedEndpointAddr, err) } return 0, err } @@ -314,7 +328,7 @@ func (p *ProxyWrapper) sendPkg(data []byte, header *PacketHeaders) error { func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn { if header.isIPv4 { - return p.wgeBPFProxy.rawConnIPv4 + return p.proxy.rawConnIPv4 } - return p.wgeBPFProxy.rawConnIPv6 + return p.proxy.rawConnIPv6 } diff --git a/client/iface/wgproxy/proxy_linux_test.go b/client/iface/wgproxy/proxy_linux_test.go index e34dd3b6b..88d4588a5 100644 --- a/client/iface/wgproxy/proxy_linux_test.go +++ b/client/iface/wgproxy/proxy_linux_test.go @@ -9,25 +9,25 @@ import ( "github.com/netbirdio/netbird/client/iface/bind" "github.com/netbirdio/netbird/client/iface/wgaddr" bindproxy "github.com/netbirdio/netbird/client/iface/wgproxy/bind" - "github.com/netbirdio/netbird/client/iface/wgproxy/ebpf" + "github.com/netbirdio/netbird/client/iface/wgproxy/loopback" "github.com/netbirdio/netbird/client/iface/wgproxy/udp" ) func seedProxies() ([]proxyInstance, error) { pl := make([]proxyInstance, 0) - ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280) - if err := ebpfProxy.Listen(); err != nil { - return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err) + loopbackProxy := loopback.NewProxy(51831, 1280) + if err := loopbackProxy.Listen(); err != nil { + return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err) } - pEbpf := proxyInstance{ - name: "ebpf kernel proxy", - proxy: ebpf.NewProxyWrapper(ebpfProxy), + pLoopback := proxyInstance{ + name: "loopback kernel proxy", + proxy: loopback.NewProxyWrapper(loopbackProxy), wgPort: 51831, - closeFn: ebpfProxy.Free, + closeFn: loopbackProxy.Free, } - pl = append(pl, pEbpf) + pl = append(pl, pLoopback) pUDP := proxyInstance{ name: "udp kernel proxy", @@ -42,18 +42,18 @@ func seedProxies() ([]proxyInstance, error) { func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) { pl := make([]proxyInstance, 0) - ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280) - if err := ebpfProxy.Listen(); err != nil { - return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err) + loopbackProxy := loopback.NewProxy(51831, 1280) + if err := loopbackProxy.Listen(); err != nil { + return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err) } - pEbpf := proxyInstance{ - name: "ebpf kernel proxy", - proxy: ebpf.NewProxyWrapper(ebpfProxy), + pLoopback := proxyInstance{ + name: "loopback kernel proxy", + proxy: loopback.NewProxyWrapper(loopbackProxy), wgPort: 51831, - closeFn: ebpfProxy.Free, + closeFn: loopbackProxy.Free, } - pl = append(pl, pEbpf) + pl = append(pl, pLoopback) pUDP := proxyInstance{ name: "udp kernel proxy", diff --git a/client/iface/wgproxy/redirect_test.go b/client/iface/wgproxy/redirect_test.go index f0d59cc64..47f571f2b 100644 --- a/client/iface/wgproxy/redirect_test.go +++ b/client/iface/wgproxy/redirect_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - "github.com/netbirdio/netbird/client/iface/wgproxy/ebpf" + "github.com/netbirdio/netbird/client/iface/wgproxy/loopback" "github.com/netbirdio/netbird/client/iface/wgproxy/udp" ) @@ -198,20 +198,20 @@ func testRedirectAs(t *testing.T, proxy Proxy, wgPort int, nbAddr, p2pEndpoint * } } -// TestRedirectAs_eBPF_IPv4 tests RedirectAs with eBPF proxy using IPv4 addresses -func TestRedirectAs_eBPF_IPv4(t *testing.T) { +// TestRedirectAs_Loopback_IPv4 tests RedirectAs with the loopback proxy using IPv4 addresses +func TestRedirectAs_Loopback_IPv4(t *testing.T) { wgPort := 51850 - ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280) - if err := ebpfProxy.Listen(); err != nil { - t.Fatalf("failed to initialize ebpf proxy: %v", err) + loopbackProxy := loopback.NewProxy(wgPort, 1280) + if err := loopbackProxy.Listen(); err != nil { + t.Fatalf("failed to initialize loopback proxy: %v", err) } defer func() { - if err := ebpfProxy.Free(); err != nil { - t.Errorf("failed to free ebpf proxy: %v", err) + if err := loopbackProxy.Free(); err != nil { + t.Errorf("failed to free loopback proxy: %v", err) } }() - proxy := ebpf.NewProxyWrapper(ebpfProxy) + proxy := loopback.NewProxyWrapper(loopbackProxy) // NetBird UDP address of the remote peer nbAddr := &net.UDPAddr{ @@ -227,20 +227,20 @@ func TestRedirectAs_eBPF_IPv4(t *testing.T) { testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint) } -// TestRedirectAs_eBPF_IPv6 tests RedirectAs with eBPF proxy using IPv6 addresses -func TestRedirectAs_eBPF_IPv6(t *testing.T) { +// TestRedirectAs_Loopback_IPv6 tests RedirectAs with the loopback proxy using IPv6 addresses +func TestRedirectAs_Loopback_IPv6(t *testing.T) { wgPort := 51851 - ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280) - if err := ebpfProxy.Listen(); err != nil { - t.Fatalf("failed to initialize ebpf proxy: %v", err) + loopbackProxy := loopback.NewProxy(wgPort, 1280) + if err := loopbackProxy.Listen(); err != nil { + t.Fatalf("failed to initialize loopback proxy: %v", err) } defer func() { - if err := ebpfProxy.Free(); err != nil { - t.Errorf("failed to free ebpf proxy: %v", err) + if err := loopbackProxy.Free(); err != nil { + t.Errorf("failed to free loopback proxy: %v", err) } }() - proxy := ebpf.NewProxyWrapper(ebpfProxy) + proxy := loopback.NewProxyWrapper(loopbackProxy) // NetBird UDP address of the remote peer nbAddr := &net.UDPAddr{ @@ -259,17 +259,17 @@ func TestRedirectAs_eBPF_IPv6(t *testing.T) { // TestRedirectAs_Multiple_Switches tests switching between multiple endpoints func TestRedirectAs_Multiple_Switches(t *testing.T) { wgPort := 51856 - ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280) - if err := ebpfProxy.Listen(); err != nil { - t.Fatalf("failed to initialize ebpf proxy: %v", err) + loopbackProxy := loopback.NewProxy(wgPort, 1280) + if err := loopbackProxy.Listen(); err != nil { + t.Fatalf("failed to initialize loopback proxy: %v", err) } defer func() { - if err := ebpfProxy.Free(); err != nil { - t.Errorf("failed to free ebpf proxy: %v", err) + if err := loopbackProxy.Free(); err != nil { + t.Errorf("failed to free loopback proxy: %v", err) } }() - proxy := ebpf.NewProxyWrapper(ebpfProxy) + proxy := loopback.NewProxyWrapper(loopbackProxy) ctx := context.Background() diff --git a/client/internal/auth/sessionwatch/watcher.go b/client/internal/auth/sessionwatch/watcher.go index e685c28d0..496903044 100644 --- a/client/internal/auth/sessionwatch/watcher.go +++ b/client/internal/auth/sessionwatch/watcher.go @@ -90,8 +90,9 @@ type StatusRecorder interface { // fallback T-FinalWarningLead dialog (suppressed when the user dismissed // the first one for the same deadline). Safe for concurrent use. type Watcher struct { - lead time.Duration - finalLead time.Duration + lead time.Duration + finalLead time.Duration + deadlineOnly bool mu sync.Mutex current time.Time @@ -102,6 +103,7 @@ type Watcher struct { dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal closed bool recorder StatusRecorder + nowFn func() time.Time } // New returns a watcher with the package defaults WarningLead and @@ -122,9 +124,17 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher { lead: lead, finalLead: final, recorder: recorder, + nowFn: time.Now, } } +// NewDeadlineOnly returns a watcher that validates and records deadlines but arms no warning timers. +func NewDeadlineOnly(recorder StatusRecorder) *Watcher { + w := New(recorder) + w.deadlineOnly = true + return w +} + // Update sets the latest deadline. Pass the zero time to clear (e.g. when // a Sync push from the server omits the field because login expiration // was disabled). @@ -181,7 +191,7 @@ func (w *Watcher) Update(deadline time.Time) error { w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} - if deadline.After(now) { + if deadline.After(now) && !w.deadlineOnly { w.armTimerLocked(deadline) } recorder := w.recorder @@ -303,6 +313,11 @@ func (w *Watcher) fire(armedFor time.Time) { w.mu.Unlock() return } + now := w.nowFn() + if isLate(now, armedFor, max(w.finalLead, 0)) { + w.fireLateLocked(armedFor, now) + return + } w.firedAt = armedFor recorder := w.recorder w.mu.Unlock() @@ -331,6 +346,14 @@ func (w *Watcher) fireFinal(armedFor time.Time) { log.Infof("auth session final-warning skipped (dismissed by user)") return } + now := w.nowFn() + if isLate(now, armedFor, 0) { + w.finalFiredAt = armedFor + w.mu.Unlock() + log.Infof("auth session final-warning skipped for deadline %s (passed %s ago)", + armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second)) + return + } w.finalFiredAt = armedFor recorder := w.recorder w.mu.Unlock() @@ -341,6 +364,39 @@ func (w *Watcher) fireFinal(armedFor time.Time) { publishWarning(recorder, armedFor, true) } +// fireLateLocked handles a T-WarningLead callback that fired inside the +// final-warning window: it sends the final warning in its place while the +// deadline has not passed and the user has not dismissed it, so a resume +// with time left still warns. The caller must hold w.mu; this helper +// releases it. +func (w *Watcher) fireLateLocked(armedFor, now time.Time) { + w.firedAt = armedFor + switch { + case w.dismissedAt.Equal(armedFor): + w.mu.Unlock() + log.Infof("auth session expiry soon warning skipped (dismissed by user)") + return + case w.finalFiredAt.Equal(armedFor): + w.mu.Unlock() + log.Infof("auth session expiry soon warning skipped (final warning already fired)") + return + case isLate(now, armedFor, 0): + w.mu.Unlock() + log.Infof("auth session expiry soon warning skipped for deadline %s (passed %s ago)", + armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second)) + return + } + w.finalFiredAt = armedFor + recorder := w.recorder + w.mu.Unlock() + if recorder == nil { + return + } + log.Infof("auth session expiry soon warning fired inside the final-warning window, sending final warning for deadline %s", + armedFor.Format(time.RFC3339)) + publishWarning(recorder, armedFor, true) +} + // armOneShotLocked schedules cb at fireAt. When fireAt is already in the // past it dispatches on the next scheduler tick so a state-change recorder // notification (invoked after w.mu is released) lands first. Caller must @@ -380,3 +436,11 @@ func publishWarning(recorder StatusRecorder, deadline time.Time, final bool) { meta, ) } + +// isLate reports whether the wall clock now has already reached armedFor +// minus cutoffLead. The timers run on the monotonic clock, which can stall +// while the host sleeps, so a timer can fire long after the window it was +// armed for. +func isLate(now, armedFor time.Time, cutoffLead time.Duration) bool { + return !now.Round(0).Before(armedFor.Add(-cutoffLead).Round(0)) +} diff --git a/client/internal/auth/sessionwatch/watcher_test.go b/client/internal/auth/sessionwatch/watcher_test.go index 4b49a94b6..cb2800978 100644 --- a/client/internal/auth/sessionwatch/watcher_test.go +++ b/client/internal/auth/sessionwatch/watcher_test.go @@ -527,3 +527,201 @@ func TestDismissBeforeUpdateIsNoop(t *testing.T) { } t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot()) } + +func TestIsLate(t *testing.T) { + armedFor := time.Date(2026, 10, 1, 12, 0, 0, 0, time.UTC) + lead := 2 * time.Minute + tests := []struct { + name string + now time.Time + cutoffLead time.Duration + want bool + }{ + {"before cutoff", armedFor.Add(-3 * time.Minute), lead, false}, + {"at cutoff", armedFor.Add(-lead), lead, true}, + {"after cutoff", armedFor.Add(-time.Minute), lead, true}, + {"zero lead before deadline", armedFor.Add(-time.Second), 0, false}, + {"zero lead at deadline", armedFor, 0, true}, + {"zero lead after deadline", armedFor.Add(time.Second), 0, true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isLate(tt.now, armedFor, tt.cutoffLead); got != tt.want { + t.Fatalf("isLate(%s, %s, %s) = %v, want %v", tt.now, armedFor, tt.cutoffLead, got, tt.want) + } + }) + } +} + +func TestIsLateIgnoresMonotonicReading(t *testing.T) { + now := time.Now() + wallOnly := now.Round(0) + if isLate(now, wallOnly.Add(time.Second), 0) { + t.Fatalf("now with monotonic reading must compare as wall clock before a later wall-only deadline") + } + if !isLate(now, wallOnly, 0) { + t.Fatalf("now with monotonic reading must compare as wall clock at an equal wall-only deadline") + } +} + +func TestLateTimerFiring(t *testing.T) { + tests := []struct { + name string + final bool + beforeDl time.Duration + wantWarns int + wantFinals int + }{ + {"warning on resume inside window", false, 3 * time.Minute, 1, 0}, + {"warning promoted to final inside final window", false, time.Minute, 0, 1}, + {"warning skipped past deadline", false, -time.Minute, 0, 0}, + {"final on resume before deadline", true, time.Minute, 0, 1}, + {"final skipped past deadline", true, -time.Minute, 0, 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + // The deadline is an hour out so the real timers never fire + // during the test; the late callback is invoked directly with an + // injected clock that simulates a resume near the deadline. + d := time.Now().Add(time.Hour).Round(0) + w.nowFn = func() time.Time { return d.Add(-tt.beforeDl) } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + if tt.final { + w.fireFinal(d) + } else { + w.fire(d) + } + + events := r.snapshot() + if got := countWhere(events, event.isWarning); got != tt.wantWarns { + t.Fatalf("expected %d warning publishes, got %d: %+v", tt.wantWarns, got, events) + } + if got := countWhere(events, event.isFinalWarning); got != tt.wantFinals { + t.Fatalf("expected %d final-warning publishes, got %d: %+v", tt.wantFinals, got, events) + } + }) + } +} + +func TestPromotedFinalWarningIsNotRepeated(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + d := time.Now().Add(time.Hour).Round(0) + now := d.Add(-time.Minute) + w.nowFn = func() time.Time { return now } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + w.fire(d) + // The final timer was suspended too, so it fires even later than the + // warning timer, here still just before the deadline. + now = d.Add(-30 * time.Second) + w.fireFinal(d) + + events := r.snapshot() + if got := countWhere(events, event.isFinalWarning); got != 1 { + t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events) + } + if got := countWhere(events, event.isWarning); got != 0 { + t.Fatalf("expected no regular warning publish, got %d: %+v", got, events) + } +} + +func TestPromotionRespectsDismiss(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + d := time.Now().Add(time.Hour).Round(0) + w.nowFn = func() time.Time { return d.Add(-time.Minute) } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + w.Dismiss() + w.fire(d) + + events := r.snapshot() + if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 { + t.Fatalf("expected no publish after dismiss, got %d: %+v", got, events) + } +} + +func TestPromotionSkippedWhenFinalAlreadyFired(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + // Both timers fall in the past after a long suspend and are dispatched + // with a zero delay, so the final callback can run before the warning one. + d := time.Now().Add(time.Hour).Round(0) + w.nowFn = func() time.Time { return d.Add(-time.Minute) } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + w.fireFinal(d) + w.fire(d) + + events := r.snapshot() + if got := countWhere(events, event.isFinalWarning); got != 1 { + t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events) + } + if got := countWhere(events, event.isWarning); got != 0 { + t.Fatalf("expected no regular warning publish, got %d: %+v", got, events) + } +} + +func TestDeadlineOnlyRecordsDeadlineWithoutWarnings(t *testing.T) { + r := &fakeRecorder{} + w := NewDeadlineOnly(r) + defer w.Close() + + // With the default leads this deadline would otherwise fire both + // timers on the next tick. + d := time.Now().Add(50 * time.Millisecond).Round(0) + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + if got := r.deadline(); !got.Equal(d) { + t.Fatalf("expected recorder deadline %v, got %v", d, got) + } + + time.Sleep(100 * time.Millisecond) + + events := r.snapshot() + if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 { + t.Fatalf("expected no publish in deadline-only mode, got %d: %+v", got, events) + } + if w.timer != nil || w.finalTimer != nil { + t.Fatal("expected no timers armed in deadline-only mode") + } +} + +func TestDeadlineOnlyStillRejectsOutOfRangeDeadlines(t *testing.T) { + r := &fakeRecorder{} + w := NewDeadlineOnly(r) + defer w.Close() + + if err := w.Update(time.Now().Add(time.Hour)); err != nil { + t.Fatalf("Update: %v", err) + } + + err := w.Update(time.Now().Add(-maxPastHorizon - time.Hour)) + if !errors.Is(err, ErrDeadlineInPast) { + t.Fatalf("expected ErrDeadlineInPast, got %v", err) + } + if got := r.deadline(); !got.IsZero() { + t.Fatalf("expected recorder cleared after rejection, got %v", got) + } +} diff --git a/client/internal/daemonaddr/grpc.go b/client/internal/daemonaddr/grpc.go new file mode 100644 index 000000000..5f0cc10cc --- /dev/null +++ b/client/internal/daemonaddr/grpc.go @@ -0,0 +1,42 @@ +package daemonaddr + +import ( + "os" + "strconv" + + log "github.com/sirupsen/logrus" +) + +const ( + // EnvMaxRecvMsgSize overrides the default gRPC max receive message size for + // connections to the daemon. Value is in bytes. + EnvMaxRecvMsgSize = "NB_DAEMON_GRPC_MAX_MSG_SIZE" + + // defaultMaxRecvMsgSize is the max gRPC receive message size used for daemon + // connections when EnvMaxRecvMsgSize is unset or invalid. It overrides the + // gRPC library default of 4 MB, which a detailed status already exceeds on a + // network of a few thousand peers. + defaultMaxRecvMsgSize = 1024 * 1024 * 16 +) + +// MaxRecvMsgSize returns the max gRPC receive message size for daemon connections +// from the environment, or defaultMaxRecvMsgSize (16 MB) if unset or invalid. +func MaxRecvMsgSize() int { + val := os.Getenv(EnvMaxRecvMsgSize) + if val == "" { + return defaultMaxRecvMsgSize + } + + size, err := strconv.Atoi(val) + if err != nil { + log.Warnf("invalid %s value %q, using default: %v", EnvMaxRecvMsgSize, val, err) + return defaultMaxRecvMsgSize + } + + if size <= 0 { + log.Warnf("invalid %s value %d, must be positive, using default", EnvMaxRecvMsgSize, size) + return defaultMaxRecvMsgSize + } + + return size +} diff --git a/client/internal/daemonaddr/grpc_test.go b/client/internal/daemonaddr/grpc_test.go new file mode 100644 index 000000000..7c4909a42 --- /dev/null +++ b/client/internal/daemonaddr/grpc_test.go @@ -0,0 +1,112 @@ +package daemonaddr + +import ( + "context" + "net" + "os" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/proto" +) + +func TestMaxRecvMsgSize(t *testing.T) { + tests := []struct { + name string + envValue string + expected int + }{ + {name: "unset returns default", envValue: "", expected: defaultMaxRecvMsgSize}, + {name: "non-numeric returns default", envValue: "abc", expected: defaultMaxRecvMsgSize}, + {name: "negative returns default", envValue: "-1", expected: defaultMaxRecvMsgSize}, + {name: "zero returns default", envValue: "0", expected: defaultMaxRecvMsgSize}, + {name: "valid value is used", envValue: "33554432", expected: 33554432}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + // Set first so the previous value is restored on cleanup, then unset to + // exercise the absent case. + t.Setenv(EnvMaxRecvMsgSize, tc.envValue) + if tc.envValue == "" { + require.NoError(t, os.Unsetenv(EnvMaxRecvMsgSize), "unset the override") + } + + assert.Equal(t, tc.expected, MaxRecvMsgSize(), "max receive message size") + }) + } +} + +// bigStatusServer answers Status with a response larger than gRPC's 4 MB default +// receive limit, which is what a detailed status on a large network looks like. +type bigStatusServer struct { + proto.UnimplementedDaemonServiceServer + payload string +} + +func (s *bigStatusServer) Status(context.Context, *proto.StatusRequest) (*proto.StatusResponse, error) { + return &proto.StatusResponse{Status: s.payload}, nil +} + +func startBigStatusServer(t *testing.T, payload string) string { + t.Helper() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err, "listen on loopback") + + srv := grpc.NewServer() + proto.RegisterDaemonServiceServer(srv, &bigStatusServer{payload: payload}) + go func() { + _ = srv.Serve(listener) + }() + t.Cleanup(srv.Stop) + + return "tcp://" + listener.Addr().String() +} + +func TestDialTargetAcceptsAStatusOverTheGrpcDefault(t *testing.T) { + payload := strings.Repeat("x", 5*1024*1024) + addr := startBigStatusServer(t, payload) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + target, opts := DialTarget(addr) + conn, err := grpc.NewClient(target, opts...) + require.NoError(t, err, "dial the daemon") + t.Cleanup(func() { _ = conn.Close() }) + + resp, err := proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{}) + require.NoError(t, err, "a detailed status must not be rejected for its size") + assert.Len(t, resp.GetStatus(), len(payload), "the whole response must arrive") +} + +// TestDialTargetRaisesTheDefaultLimit is the negative control: the same response +// over a connection carrying gRPC's own defaults is refused, which is the failure +// reported by `netbird status -d` on a large deployment. +func TestDialTargetRaisesTheDefaultLimit(t *testing.T) { + payload := strings.Repeat("x", 5*1024*1024) + addr := startBigStatusServer(t, payload) + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + conn, err := grpc.NewClient( + strings.TrimPrefix(addr, "tcp://"), + grpc.WithTransportCredentials(insecure.NewCredentials()), + ) + require.NoError(t, err, "dial with the library defaults") + t.Cleanup(func() { _ = conn.Close() }) + + _, err = proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{}) + require.Error(t, err, "the library default must reject this response") + assert.Equal(t, codes.ResourceExhausted, status.Code(err), "gRPC rejects an oversized message") +} diff --git a/client/internal/daemonaddr/pipe.go b/client/internal/daemonaddr/pipe.go index 51815ef5e..bf1c8fdd0 100644 --- a/client/internal/daemonaddr/pipe.go +++ b/client/internal/daemonaddr/pipe.go @@ -36,7 +36,10 @@ const ( // address. The npipe scheme needs a context dialer because gRPC has no // named-pipe resolver; unix and tcp are handled by gRPC itself. func DialTarget(addr string) (string, []grpc.DialOption) { - opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())} + opts := []grpc.DialOption{ + grpc.WithTransportCredentials(insecure.NewCredentials()), + grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(MaxRecvMsgSize())), + } if name, ok := strings.CutPrefix(addr, pipeScheme); ok { paths := PipePaths(name) diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go index b362ae293..f4d1c3598 100644 --- a/client/internal/debug/debug.go +++ b/client/internal/debug/debug.go @@ -379,9 +379,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen } } +// bundleFilePattern names the bundle zips Generate creates in tempDir; the +// asterisk is filled in by os.CreateTemp. +const bundleFilePattern = "netbird.debug.*.zip" + +const exportedBundlePrefix = "netbird.debug-file." + +const exportedBundleMaxAge = 24 * time.Hour + +// RemoveStaleBundles deletes bundle zips that an interrupted generation or +// upload left behind in dir. Only files older than maxAge go, so a bundle that +// another caller is still writing or uploading in the same directory survives. +// Exported bundles are kept for exportedBundleMaxAge instead. +func RemoveStaleBundles(dir string, maxAge time.Duration) { + removeStaleFiles(dir, bundleFilePattern, maxAge) + removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge) +} + +// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern +// and returns the new path. The caller owns the file from then on; an export +// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles. +func ExportBundle(path string) (string, error) { + base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0]) + exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base) + if err := os.Rename(path, exported); err != nil { + return "", fmt.Errorf("export debug bundle: %w", err) + } + return exported, nil +} + // Generate creates a debug bundle and returns the location. func (g *BundleGenerator) Generate() (resp string, err error) { - bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip") + bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern) if err != nil { return "", fmt.Errorf("create zip file: %w", err) } @@ -1725,3 +1754,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any { } return v } + +func removeStaleFiles(dir, pattern string, maxAge time.Duration) { + matches, err := filepath.Glob(filepath.Join(dir, pattern)) + if err != nil { + log.Debugf("glob stale debug bundles in %s: %v", dir, err) + return + } + + cutoff := time.Now().Add(-maxAge) + for _, path := range matches { + info, err := os.Stat(path) + if err != nil || info.ModTime().After(cutoff) { + continue + } + if err := os.Remove(path); err != nil { + if !errors.Is(err, fs.ErrNotExist) { + log.Warnf("remove stale debug bundle %s: %v", path, err) + } + continue + } + log.Infof("removed stale debug bundle %s", path) + } +} diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go index 17d520358..6a810bccc 100644 --- a/client/internal/debug/debug_test.go +++ b/client/internal/debug/debug_test.go @@ -4,6 +4,7 @@ import ( "archive/zip" "bytes" "encoding/json" + "fmt" "net" "net/netip" "net/url" @@ -969,3 +970,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string { func newAnonymizerForTest() *anonymize.Anonymizer { return anonymize.NewAnonymizer(anonymize.DefaultAddresses()) } + +func TestRemoveStaleBundles(t *testing.T) { + dir := t.TempDir() + stale := filepath.Join(dir, "netbird.debug.111.zip") + fresh := filepath.Join(dir, "netbird.debug.222.zip") + other := filepath.Join(dir, "netbird.debug.333.txt") + owned := filepath.Join(dir, "netbird.debug.444.zip") + abandoned := filepath.Join(dir, "netbird.debug.555.zip") + for _, p := range []string{stale, fresh, other, owned, abandoned} { + require.NoError(t, os.WriteFile(p, []byte("x"), 0o600)) + } + exported, err := ExportBundle(owned) + require.NoError(t, err) + exportedAbandoned, err := ExportBundle(abandoned) + require.NoError(t, err) + old := time.Now().Add(-2 * time.Hour) + for _, p := range []string{stale, other, exported} { + require.NoError(t, os.Chtimes(p, old, old)) + } + ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour) + require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient)) + + RemoveStaleBundles(dir, time.Hour) + + assert.NoFileExists(t, stale, "bundle older than maxAge should be removed") + assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading") + assert.FileExists(t, other, "files outside the bundle pattern must not be touched") + assert.NoFileExists(t, owned) + assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge") + assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned") +} + +func TestBundleIncludesNetworkMap(t *testing.T) { + for _, anonymize := range []bool{false, true} { + t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) { + g := NewBundleGenerator(GeneratorDependencies{ + SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}}, + }, BundleConfig{Anonymize: anonymize}) + + require.Contains(t, bundleEntries(t, g), "network_map.json") + }) + } +} + +func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) { + g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{}) + + require.NotContains(t, bundleEntries(t, g), "network_map.json") +} diff --git a/client/internal/dns/host_windows.go b/client/internal/dns/host_windows.go index 948000a3d..6462d0c37 100644 --- a/client/internal/dns/host_windows.go +++ b/client/internal/dns/host_windows.go @@ -124,19 +124,9 @@ func newHostManager(wgInterface WGIface) (*registryConfigurator, error) { return nil, err } - var useGPO bool - k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE) - if err != nil { - log.Debugf("failed to open GPO DNS policy root: %v", err) - } else { - closer(k) - useGPO = true - log.Infof("detected GPO DNS policy configuration, using policy store") - } - configurator := ®istryConfigurator{ guid: guid, - gpo: useGPO, + gpo: useGPOPolicyStore(), } origNameservers, err := configurator.captureOriginalNameservers() @@ -576,14 +566,22 @@ func (r *registryConfigurator) setInterfaceRegistryKeyStringValue(key, value str return nil } +// deleteInterfaceRegistryKeyProperty removes a value from the interface key. +// A value that is already gone, or an interface key that is, is not an error: +// the caller asked for the value not to be there, and a cleanup that runs twice +// has to reach its later steps on the second run as well. func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey string) error { regKey, err := r.getInterfaceRegistryKey() - if err != nil { + switch { + case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND): + log.Debugf("interface key of %s does not exist, nothing to delete %s from", r.guid, propertyKey) + return nil + case err != nil: return fmt.Errorf("get interface registry key: %w", err) } defer closer(regKey) - if err := regKey.DeleteValue(propertyKey); err != nil { + if err := regKey.DeleteValue(propertyKey); err != nil && !errors.Is(err, registry.ErrNotExist) { return fmt.Errorf("delete registry key %s: %w", propertyKey, err) } return nil @@ -612,7 +610,12 @@ func (r *registryConfigurator) restoreHostDNS() error { go r.flushDNSCache() - return nil + // Last, and only on the way out, once no rule of ours is left: during a + // session the store is where the rules of this run live, and emptying it + // mid-session would have the next rule recreate it anyway. Propagated so a + // failure keeps the shutdown state for the next run to retry, rather than + // leaving the store to hold up every rule change from here on. + return removeEmptyGPOPolicyStore() } // removeDNSMatchPolicies deletes every NRPT rule this client may have created, @@ -651,6 +654,73 @@ func (r *registryConfigurator) restoreUncleanShutdownDNS() error { return r.restoreHostDNS() } +// useGPOPolicyStore reports whether NRPT rules have to go into the group policy +// store, and clears an empty one out of the way first. +// +// The order is the point. A store left empty by an earlier run would otherwise +// decide this run too, sending its rules somewhere the resolver only reads when +// the policy engine next applies DNS client policy. Removing it before the +// choice is made leaves the local store authoritative for the whole session, +// including the first one after an upgrade. +func useGPOPolicyStore() bool { + if err := removeEmptyGPOPolicyStore(); err != nil { + // Nothing to retry against here: the worst case is the run going + // through the group policy store, which is where it would have gone + // before this check existed. + log.Warnf("%v", err) + } + + k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE) + if err != nil { + log.Debugf("failed to open GPO DNS policy root: %v", err) + return false + } + closer(k) + + log.Infof("detected GPO DNS policy configuration, using policy store") + return true +} + +// removeEmptyGPOPolicyStore deletes the group policy DnsPolicyConfig key once +// nothing is left in it. The key survives the deletion of the last rule it +// held, and the client treats its presence as "group policy configures the +// NRPT", so an empty one left behind keeps every later run writing rules there. +// Rules in that store reach the resolver only when the policy engine next +// applies DNS client policy, and a rule this client writes belongs to no GPO, +// so nothing schedules that application: both adding and removing a rule are +// held up by a minute or more, and for a removal that is a catch-all rule +// resolving every name over an interface that no longer exists. With the store +// absent the local one is authoritative and a change applies at once. +// +// A store that still holds rules, values or subkeys of somebody else's is left +// alone. +func removeEmptyGPOPolicyStore() error { + k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE) + switch { + case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND): + return nil + case err != nil: + return fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err) + } + + info, err := k.Stat() + closer(k) + if err != nil { + return fmt.Errorf("stat HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err) + } + + if info.SubKeyCount != 0 || info.ValueCount != 0 { + return nil + } + + if err := registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot); err != nil { + return fmt.Errorf("delete empty HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err) + } + + log.Infof("removed the empty GPO DNS policy store, leaving the local one authoritative") + return nil +} + // listNRPTRuleKeys returns the names of our NRPT rule keys under a policy store // root. An absent root holds nothing to clean up, which is the normal state of // the GPO store on a machine without DNS Client policy. diff --git a/client/internal/dns/host_windows_test.go b/client/internal/dns/host_windows_test.go index 7aef64590..353f6adbc 100644 --- a/client/internal/dns/host_windows_test.go +++ b/client/internal/dns/host_windows_test.go @@ -8,6 +8,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "golang.org/x/sys/windows/registry" + + "github.com/netbirdio/netbird/client/internal/winregistry" ) // TestNRPTEntriesCleanupOnConfigChange tests that old NRPT entries are properly cleaned up @@ -405,3 +407,130 @@ func TestNRPTDomainBatching(t *testing.T) { }) } } + +// TestRemoveEmptyGPOPolicyStore verifies that cleanup takes the GPO policy +// store itself with it once our rules are gone, since the store existing keeps +// the local one from being applied, and that a store with somebody else's rule +// in it is left alone. +func TestRemoveEmptyGPOPolicyStore(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + t.Cleanup(func() { cleanupRegistryKeys(t) }) + cleanupRegistryKeys(t) + + testIP := netip.MustParseAddr("100.64.0.1") + cfg := ®istryConfigurator{gpo: true} + + // a store holding a rule of ours is kept, because the rule is still applied + require.NoError(t, cfg.addDNSMatchPolicy([]string{".example.com"}, testIP)) + exists, err := registryKeyExists(gpoDnsPolicyConfigMatchPath + "-0") + require.NoError(t, err) + require.True(t, exists, "Should write the rule to the GPO policy store") + + require.NoError(t, removeEmptyGPOPolicyStore()) + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.True(t, exists, "Should keep a policy store that still holds a rule") + + // once the rules are gone the store goes with them + require.NoError(t, cfg.removeDNSMatchPolicies()) + require.NoError(t, removeEmptyGPOPolicyStore()) + + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.False(t, exists, "Should remove the GPO policy store once it is empty") + + // A store is not ours to remove while somebody else has a rule in it. The + // rule is written volatile like our own: the rules above created the parent + // chain volatile, and Windows refuses a stable subkey under a volatile + // parent. + foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}` + foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE) + require.NoError(t, err, "Should create a foreign GPO rule") + foreignKey.Close() + t.Cleanup(func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule) + _ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot) + }) + + require.NoError(t, cfg.removeDNSMatchPolicies()) + require.NoError(t, removeEmptyGPOPolicyStore()) + + exists, err = registryKeyExists(foreignRule) + require.NoError(t, err) + assert.True(t, exists, "Should not remove a foreign rule") + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.True(t, exists, "Should keep a policy store that still holds a foreign rule") +} + +// TestDeleteInterfaceRegistryKeyPropertyTwice verifies that removing a value +// that is already gone, or one on an interface key that is, reports success. +// Teardown runs again after a failed cleanup, and the steps that follow this +// one have to be reached on that second run. +func TestDeleteInterfaceRegistryKeyPropertyTwice(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + testGUID := "{12345678-1234-1234-1234-123456789ABC}" + interfacePath := InterfaceConfigPath + `\` + testGUID + testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE) + require.NoError(t, err, "Should create test interface registry key") + testKey.Close() + t.Cleanup(func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath) + }) + + cfg := ®istryConfigurator{guid: testGUID} + + require.NoError(t, cfg.setInterfaceRegistryKeyStringValue(interfaceConfigSearchListKey, "example.com")) + require.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey)) + assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey), + "Should report success for a value that is already gone") + + // and with the interface key itself gone, as it is once the adapter is + require.NoError(t, registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath)) + assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey), + "Should report success when the interface key does not exist") +} + +// TestUseGPOPolicyStoreClearsEmptyStore verifies that the store is cleared +// before it is consulted, so an empty one left by an earlier run does not send +// this run's rules to the group policy store. A store somebody else has a rule +// in still decides where the rules go. +func TestUseGPOPolicyStoreClearsEmptyStore(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + t.Cleanup(func() { cleanupRegistryKeys(t) }) + cleanupRegistryKeys(t) + + // the leftover an earlier run used to keep, which the client read as + // "group policy configures the NRPT" for every run after it + emptyStore, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.SET_VALUE) + require.NoError(t, err, "Should create the GPO policy store") + emptyStore.Close() + + assert.False(t, useGPOPolicyStore(), "An empty store should not decide where the rules go") + exists, err := registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.False(t, exists, "Should clear the empty store before consulting it") + + foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}` + foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE) + require.NoError(t, err, "Should create a foreign GPO rule") + foreignKey.Close() + t.Cleanup(func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule) + _ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot) + }) + + assert.True(t, useGPOPolicyStore(), "A store holding a rule should decide where the rules go") + exists, err = registryKeyExists(GPODNSPolicyConfigRoot) + require.NoError(t, err) + assert.True(t, exists, "Should keep a store that holds a rule") +} diff --git a/client/internal/dns/server_privileged_test.go b/client/internal/dns/server_privileged_test.go index a17044cf5..270e3bf91 100644 --- a/client/internal/dns/server_privileged_test.go +++ b/client/internal/dns/server_privileged_test.go @@ -9,9 +9,9 @@ import ( "os" "testing" - "go.uber.org/mock/gomock" "github.com/miekg/dns" "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "github.com/netbirdio/netbird/client/iface" @@ -24,6 +24,10 @@ import ( nbdns "github.com/netbirdio/netbird/dns" ) +// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist +// carries. Declared here rather than imported because profilemanager imports this package. +var testIFaceBlackList = []string{"wt", "utun", "tun0"} + func TestUpdateDNSServer(t *testing.T) { nameServers := []nbdns.NameServer{ @@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { privKey, _ := wgtypes.GenerateKey() - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun230%d", n), @@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) t.Setenv("NB_WG_KERNEL_DISABLED", "true") - newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) - if err != nil { - t.Errorf("create stdnet: %v", err) - return - } + newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}) privKey, _ := wgtypes.GeneratePrivateKey() opts := iface.WGIFaceOpts{ diff --git a/client/internal/dns/server_test.go b/client/internal/dns/server_test.go index 0144a4a8b..414890158 100644 --- a/client/internal/dns/server_test.go +++ b/client/internal/dns/server_test.go @@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) t.Setenv("NB_WG_KERNEL_DISABLED", "true") - newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) - if err != nil { - t.Fatalf("create stdnet: %v", err) - return nil, err - } + newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}) privKey, _ := wgtypes.GeneratePrivateKey() diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.go b/client/internal/ebpf/ebpf/bpf_bpfeb.go deleted file mode 100644 index 4b6230217..000000000 --- a/client/internal/ebpf/ebpf/bpf_bpfeb.go +++ /dev/null @@ -1,148 +0,0 @@ -// Code generated by bpf2go; DO NOT EDIT. -//go:build mips || mips64 || ppc64 || s390x - -package ebpf - -import ( - "bytes" - _ "embed" - "fmt" - "io" - - "github.com/cilium/ebpf" -) - -// loadBpf returns the embedded CollectionSpec for bpf. -func loadBpf() (*ebpf.CollectionSpec, error) { - reader := bytes.NewReader(_BpfBytes) - spec, err := ebpf.LoadCollectionSpecFromReader(reader) - if err != nil { - return nil, fmt.Errorf("can't load bpf: %w", err) - } - - return spec, err -} - -// loadBpfObjects loads bpf and converts it into a struct. -// -// The following types are suitable as obj argument: -// -// *bpfObjects -// *bpfPrograms -// *bpfMaps -// -// See ebpf.CollectionSpec.LoadAndAssign documentation for details. -func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error { - spec, err := loadBpf() - if err != nil { - return err - } - - return spec.LoadAndAssign(obj, opts) -} - -// bpfSpecs contains maps and programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfSpecs struct { - bpfProgramSpecs - bpfMapSpecs - bpfVariableSpecs -} - -// bpfProgramSpecs contains programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfProgramSpecs struct { - NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"` -} - -// bpfMapSpecs contains maps before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfMapSpecs struct { - NbFeatures *ebpf.MapSpec `ebpf:"nb_features"` - NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"` -} - -// bpfVariableSpecs contains global variables before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfVariableSpecs struct { - FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"` - MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"` - MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"` - MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"` - ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"` - WgPort *ebpf.VariableSpec `ebpf:"wg_port"` -} - -// bpfObjects contains all objects after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfObjects struct { - bpfPrograms - bpfMaps - bpfVariables -} - -func (o *bpfObjects) Close() error { - return _BpfClose( - &o.bpfPrograms, - &o.bpfMaps, - ) -} - -// bpfMaps contains all maps after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfMaps struct { - NbFeatures *ebpf.Map `ebpf:"nb_features"` - NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"` -} - -func (m *bpfMaps) Close() error { - return _BpfClose( - m.NbFeatures, - m.NbWgProxySettingsMap, - ) -} - -// bpfVariables contains all global variables after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfVariables struct { - FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"` - MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"` - MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"` - MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"` - ProxyPort *ebpf.Variable `ebpf:"proxy_port"` - WgPort *ebpf.Variable `ebpf:"wg_port"` -} - -// bpfPrograms contains all programs after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfPrograms struct { - NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"` -} - -func (p *bpfPrograms) Close() error { - return _BpfClose( - p.NbXdpProg, - ) -} - -func _BpfClose(closers ...io.Closer) error { - for _, closer := range closers { - if err := closer.Close(); err != nil { - return err - } - } - return nil -} - -// Do not access this directly. -// -//go:embed bpf_bpfeb.o -var _BpfBytes []byte diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.o b/client/internal/ebpf/ebpf/bpf_bpfeb.o deleted file mode 100644 index b435d4964..000000000 Binary files a/client/internal/ebpf/ebpf/bpf_bpfeb.o and /dev/null differ diff --git a/client/internal/ebpf/ebpf/bpf_bpfel.go b/client/internal/ebpf/ebpf/bpf_bpfel.go deleted file mode 100644 index f56efc901..000000000 --- a/client/internal/ebpf/ebpf/bpf_bpfel.go +++ /dev/null @@ -1,148 +0,0 @@ -// Code generated by bpf2go; DO NOT EDIT. -//go:build 386 || amd64 || arm || arm64 || loong64 || mips64le || mipsle || ppc64le || riscv64 || wasm - -package ebpf - -import ( - "bytes" - _ "embed" - "fmt" - "io" - - "github.com/cilium/ebpf" -) - -// loadBpf returns the embedded CollectionSpec for bpf. -func loadBpf() (*ebpf.CollectionSpec, error) { - reader := bytes.NewReader(_BpfBytes) - spec, err := ebpf.LoadCollectionSpecFromReader(reader) - if err != nil { - return nil, fmt.Errorf("can't load bpf: %w", err) - } - - return spec, err -} - -// loadBpfObjects loads bpf and converts it into a struct. -// -// The following types are suitable as obj argument: -// -// *bpfObjects -// *bpfPrograms -// *bpfMaps -// -// See ebpf.CollectionSpec.LoadAndAssign documentation for details. -func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error { - spec, err := loadBpf() - if err != nil { - return err - } - - return spec.LoadAndAssign(obj, opts) -} - -// bpfSpecs contains maps and programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfSpecs struct { - bpfProgramSpecs - bpfMapSpecs - bpfVariableSpecs -} - -// bpfProgramSpecs contains programs before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfProgramSpecs struct { - NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"` -} - -// bpfMapSpecs contains maps before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfMapSpecs struct { - NbFeatures *ebpf.MapSpec `ebpf:"nb_features"` - NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"` -} - -// bpfVariableSpecs contains global variables before they are loaded into the kernel. -// -// It can be passed ebpf.CollectionSpec.Assign. -type bpfVariableSpecs struct { - FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"` - MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"` - MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"` - MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"` - ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"` - WgPort *ebpf.VariableSpec `ebpf:"wg_port"` -} - -// bpfObjects contains all objects after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfObjects struct { - bpfPrograms - bpfMaps - bpfVariables -} - -func (o *bpfObjects) Close() error { - return _BpfClose( - &o.bpfPrograms, - &o.bpfMaps, - ) -} - -// bpfMaps contains all maps after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfMaps struct { - NbFeatures *ebpf.Map `ebpf:"nb_features"` - NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"` -} - -func (m *bpfMaps) Close() error { - return _BpfClose( - m.NbFeatures, - m.NbWgProxySettingsMap, - ) -} - -// bpfVariables contains all global variables after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfVariables struct { - FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"` - MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"` - MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"` - MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"` - ProxyPort *ebpf.Variable `ebpf:"proxy_port"` - WgPort *ebpf.Variable `ebpf:"wg_port"` -} - -// bpfPrograms contains all programs after they have been loaded into the kernel. -// -// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. -type bpfPrograms struct { - NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"` -} - -func (p *bpfPrograms) Close() error { - return _BpfClose( - p.NbXdpProg, - ) -} - -func _BpfClose(closers ...io.Closer) error { - for _, closer := range closers { - if err := closer.Close(); err != nil { - return err - } - } - return nil -} - -// Do not access this directly. -// -//go:embed bpf_bpfel.o -var _BpfBytes []byte diff --git a/client/internal/ebpf/ebpf/bpf_bpfel.o b/client/internal/ebpf/ebpf/bpf_bpfel.o deleted file mode 100644 index a388b6d6d..000000000 Binary files a/client/internal/ebpf/ebpf/bpf_bpfel.o and /dev/null differ diff --git a/client/internal/ebpf/ebpf/manager_linux.go b/client/internal/ebpf/ebpf/manager_linux.go deleted file mode 100644 index a13f5f19a..000000000 --- a/client/internal/ebpf/ebpf/manager_linux.go +++ /dev/null @@ -1,115 +0,0 @@ -package ebpf - -import ( - _ "embed" - "net" - "sync" - - "github.com/cilium/ebpf/link" - "github.com/cilium/ebpf/rlimit" - log "github.com/sirupsen/logrus" - - "github.com/netbirdio/netbird/client/internal/ebpf/manager" -) - -const ( - mapKeyFeatures uint32 = 0 - - featureFlagWGProxy = 0b00000001 -) - -var ( - singleton manager.Manager - singletonLock = &sync.Mutex{} -) - -// required packages libbpf-dev, libc6-dev-i386-amd64-cross - -// GeneralManager is used to load multiple eBPF programs with a custom check (if then) done in prog.c -// The manager simply adds a feature (byte) of each program to a map that is shared between the userspace and kernel. -// When packet arrives, the C code checks for each feature (if it is set) and executes each enabled program (e.g., wg_proxy.c). -// -//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc clang-14 bpf src/prog.c -- -I /usr/x86_64-linux-gnu/include -include src/bpf_map_def.h -type GeneralManager struct { - lock sync.Mutex - link link.Link - featureFlags uint16 - bpfObjs bpfObjects -} - -// GetEbpfManagerInstance return a static eBpf Manager instance -func GetEbpfManagerInstance() manager.Manager { - singletonLock.Lock() - defer singletonLock.Unlock() - if singleton != nil { - return singleton - } - singleton = &GeneralManager{} - return singleton -} - -func (tf *GeneralManager) setFeatureFlag(feature uint16) { - tf.featureFlags |= feature -} - -func (tf *GeneralManager) loadXdp() error { - if tf.link != nil { - return nil - } - // it required for Docker - err := rlimit.RemoveMemlock() - if err != nil { - return err - } - - iFace, err := net.InterfaceByName("lo") - if err != nil { - return err - } - - // load pre-compiled programs into the kernel. - err = loadBpfObjects(&tf.bpfObjs, nil) - if err != nil { - return err - } - - tf.link, err = link.AttachXDP(link.XDPOptions{ - Program: tf.bpfObjs.NbXdpProg, - Interface: iFace.Index, - }) - - if err != nil { - _ = tf.bpfObjs.Close() - tf.link = nil - return err - } - return nil -} - -func (tf *GeneralManager) unsetFeatureFlag(feature uint16) error { - tf.lock.Lock() - defer tf.lock.Unlock() - tf.featureFlags &^= feature - - if tf.link == nil { - return nil - } - - if tf.featureFlags == 0 { - return tf.close() - } - - return tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags) -} - -func (tf *GeneralManager) close() error { - log.Debugf("detach ebpf program ") - err := tf.bpfObjs.Close() - if err != nil { - log.Warnf("failed to close eBpf objects: %s", err) - } - - err = tf.link.Close() - tf.link = nil - return err -} diff --git a/client/internal/ebpf/ebpf/manager_linux_test.go b/client/internal/ebpf/ebpf/manager_linux_test.go deleted file mode 100644 index e09fcb977..000000000 --- a/client/internal/ebpf/ebpf/manager_linux_test.go +++ /dev/null @@ -1,31 +0,0 @@ -package ebpf - -import ( - "testing" -) - -func TestManager_setFeatureFlag(t *testing.T) { - mgr := GeneralManager{} - mgr.setFeatureFlag(featureFlagWGProxy) - if mgr.featureFlags != featureFlagWGProxy { - t.Errorf("invalid feature state") - } - - mgr.setFeatureFlag(featureFlagWGProxy) - if mgr.featureFlags != featureFlagWGProxy { - t.Errorf("setting a flag twice must be idempotent, got: %d", mgr.featureFlags) - } -} - -func TestManager_unsetFeatureFlag(t *testing.T) { - mgr := GeneralManager{} - mgr.setFeatureFlag(featureFlagWGProxy) - - err := mgr.unsetFeatureFlag(featureFlagWGProxy) - if err != nil { - t.Errorf("unexpected error: %s", err) - } - if mgr.featureFlags != 0 { - t.Errorf("invalid feature state, expected: %d, got: %d", 0, mgr.featureFlags) - } -} diff --git a/client/internal/ebpf/ebpf/src/bpf_map_def.h b/client/internal/ebpf/ebpf/src/bpf_map_def.h deleted file mode 100644 index 9528fb592..000000000 --- a/client/internal/ebpf/ebpf/src/bpf_map_def.h +++ /dev/null @@ -1,16 +0,0 @@ -// libbpf 1.0 removed struct bpf_map_def, but the programs here keep the legacy -// map definitions: they load on kernels built without BTF, which BTF-style -// (SEC(".maps")) definitions do not. Define the struct ourselves so the -// programs compile against current libbpf headers. -#ifndef NB_BPF_MAP_DEF_H -#define NB_BPF_MAP_DEF_H - -struct bpf_map_def { - unsigned int type; - unsigned int key_size; - unsigned int value_size; - unsigned int max_entries; - unsigned int map_flags; -}; - -#endif diff --git a/client/internal/ebpf/ebpf/src/prog.c b/client/internal/ebpf/ebpf/src/prog.c deleted file mode 100644 index 44ee53458..000000000 --- a/client/internal/ebpf/ebpf/src/prog.c +++ /dev/null @@ -1,54 +0,0 @@ -#include -#include // ETH_P_IP -#include -#include -#include -#include -#include -#include "wg_proxy.c" - -const __u16 flag_feature_wg_proxy = 0b01; - -const __u32 map_key_features = 0; -struct bpf_map_def SEC("maps") nb_features = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u16), - .max_entries = 10, -}; - -SEC("xdp") -int nb_xdp_prog(struct xdp_md *ctx) { - __u16 *features; - features = bpf_map_lookup_elem(&nb_features, &map_key_features); - if (!features) { - return XDP_PASS; - } - - void *data = (void *)(long)ctx->data; - void *data_end = (void *)(long)ctx->data_end; - struct ethhdr *eth = data; - struct iphdr *ip = (data + sizeof(struct ethhdr)); - struct udphdr *udp = (data + sizeof(struct ethhdr) + sizeof(struct iphdr)); - - // return early if not enough data - if (data + sizeof(struct ethhdr) + sizeof(struct iphdr) + sizeof(struct udphdr) > data_end){ - return XDP_PASS; - } - - // skip non IPv4 packages - if (eth->h_proto != htons(ETH_P_IP)) { - return XDP_PASS; - } - - // skip non UPD packages - if (ip->protocol != IPPROTO_UDP) { - return XDP_PASS; - } - - if (*features & flag_feature_wg_proxy) { - xdp_wg_proxy(ip, udp); - } - return XDP_PASS; -} -char _license[] SEC("license") = "GPL"; diff --git a/client/internal/ebpf/ebpf/src/readme.md b/client/internal/ebpf/ebpf/src/readme.md deleted file mode 100644 index aa47847da..000000000 --- a/client/internal/ebpf/ebpf/src/readme.md +++ /dev/null @@ -1,27 +0,0 @@ -# XDP programs - -`prog.c` is attached to the `lo` device and dispatches to the features enabled in the -`nb_features` map. The only feature is the WireGuard proxy (`wg_proxy.c`): it rewrites -loopback UDP sent from the WireGuard listen port so it reaches the userspace relay proxy -port instead, and swaps the peer endpoint port into the source so the proxy can tell -peers apart. - -Maps use the legacy `struct bpf_map_def` form, defined in `bpf_map_def.h` because libbpf -1.0 removed it. They load on kernels built without BTF, which BTF-style (`SEC(".maps")`) -definitions do not. - -Regenerate the objects with `go generate ./client/internal/ebpf/ebpf/`; it needs -`clang-14`. Loading a regenerated object needs root, attaching it needs `bpf_link` -(kernel >= 5.7), and only one XDP program can own `lo` at a time. - -# Debug - -The CONFIG_BPF_EVENTS kernel module is required for bpf_printk. -Apply this code to use bpf_printk -``` -#define bpf_printk(fmt, ...) \ - ({ \ - char ____fmt[] = fmt; \ - bpf_trace_printk(____fmt, sizeof(____fmt), ##__VA_ARGS__); \ - }) -``` diff --git a/client/internal/ebpf/ebpf/src/wg_proxy.c b/client/internal/ebpf/ebpf/src/wg_proxy.c deleted file mode 100644 index 5e7474928..000000000 --- a/client/internal/ebpf/ebpf/src/wg_proxy.c +++ /dev/null @@ -1,60 +0,0 @@ -const __u32 map_key_proxy_port = 0; -const __u32 map_key_wg_port = 1; - -struct bpf_map_def SEC("maps") nb_wg_proxy_settings_map = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u16), - .max_entries = 10, -}; - -__u16 proxy_port = 0; -__u16 wg_port = 0; - -bool read_port_settings() { - __u16 *value; - value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_proxy_port); - if (!value) { - return false; - } - - proxy_port = *value; - - value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_wg_port); - if (!value) { - return false; - } - wg_port = htons(*value); - - return true; -} - -int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) { - if (proxy_port == 0 || wg_port == 0) { - if (!read_port_settings()){ - return XDP_PASS; - } - // bpf_printk("proxy port: %d, wg port: %d", proxy_port, wg_port); - } - - // 2130706433 = 127.0.0.1 - if (ip->daddr != htonl(2130706433)) { - return XDP_PASS; - } - - if (udp->source != wg_port){ - return XDP_PASS; - } - - __be16 new_src_port = udp->dest; - __be16 new_dst_port = htons(proxy_port); - udp->dest = new_dst_port; - udp->source = new_src_port; - - // The ports are covered by the UDP checksum. This is an IPv4 loopback hop - // and the payload is already integrity-protected, so clear the checksum (a - // zero UDP checksum means "not computed" for IPv4) rather than leave a - // stale value the kernel would drop as UDP_CSUM. - udp->check = 0; - return XDP_PASS; -} diff --git a/client/internal/ebpf/ebpf/wg_proxy_linux.go b/client/internal/ebpf/ebpf/wg_proxy_linux.go deleted file mode 100644 index 4e0df7329..000000000 --- a/client/internal/ebpf/ebpf/wg_proxy_linux.go +++ /dev/null @@ -1,41 +0,0 @@ -package ebpf - -import log "github.com/sirupsen/logrus" - -const ( - mapKeyProxyPort uint32 = 0 - mapKeyWgPort uint32 = 1 -) - -func (tf *GeneralManager) LoadWgProxy(proxyPort, wgPort int) error { - log.Debugf("load ebpf WG proxy") - tf.lock.Lock() - defer tf.lock.Unlock() - - err := tf.loadXdp() - if err != nil { - return err - } - - err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyProxyPort, uint16(proxyPort)) - if err != nil { - return err - } - - err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyWgPort, uint16(wgPort)) - if err != nil { - return err - } - - tf.setFeatureFlag(featureFlagWGProxy) - err = tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags) - if err != nil { - return err - } - return nil -} - -func (tf *GeneralManager) FreeWGProxy() error { - log.Debugf("free ebpf WG proxy") - return tf.unsetFeatureFlag(featureFlagWGProxy) -} diff --git a/client/internal/ebpf/instantiater_linux.go b/client/internal/ebpf/instantiater_linux.go deleted file mode 100644 index 20d8145b4..000000000 --- a/client/internal/ebpf/instantiater_linux.go +++ /dev/null @@ -1,15 +0,0 @@ -//go:build !android - -package ebpf - -import ( - "github.com/netbirdio/netbird/client/internal/ebpf/ebpf" - "github.com/netbirdio/netbird/client/internal/ebpf/manager" -) - -// GetEbpfManagerInstance is a wrapper function. This encapsulation is required because if the code import the internal -// ebpf package the Go compiler will include the object files. But it is not supported on Android. It can cause instant -// panic on older Android version. -func GetEbpfManagerInstance() manager.Manager { - return ebpf.GetEbpfManagerInstance() -} diff --git a/client/internal/ebpf/instantiater_nonlinux.go b/client/internal/ebpf/instantiater_nonlinux.go deleted file mode 100644 index b7c38733a..000000000 --- a/client/internal/ebpf/instantiater_nonlinux.go +++ /dev/null @@ -1,10 +0,0 @@ -//go:build !linux || android - -package ebpf - -import "github.com/netbirdio/netbird/client/internal/ebpf/manager" - -// GetEbpfManagerInstance return error because ebpf is not supported on all os -func GetEbpfManagerInstance() manager.Manager { - panic("unsupported os") -} diff --git a/client/internal/ebpf/manager/manager.go b/client/internal/ebpf/manager/manager.go deleted file mode 100644 index fdc5d8d82..000000000 --- a/client/internal/ebpf/manager/manager.go +++ /dev/null @@ -1,7 +0,0 @@ -package manager - -// Manager is used to load multiple eBPF programs. E.g., the WireGuard proxy -type Manager interface { - LoadWgProxy(proxyPort, wgPort int) error - FreeWGProxy() error -} diff --git a/client/internal/elevate/trusted.go b/client/internal/elevate/trusted.go index c11054c45..98e05fde5 100644 --- a/client/internal/elevate/trusted.go +++ b/client/internal/elevate/trusted.go @@ -6,6 +6,17 @@ import ( "path/filepath" ) +// CheckOnlyOwnerWritable reports an error unless path, and every directory +// leading to it, is owned by an account that can already act with the privileges +// the caller holds, and is writable by nobody else. +// +// Exported for callers outside elevation that read a file while privileged and +// then act on what it says: the same question this package asks of an +// executable, asked of a configuration file. +func CheckOnlyOwnerWritable(path string) error { + return checkOnlyOwnerWritable(path) +} + // trustedSelf returns the path of this executable, provided it is one we are // willing to have run as root. // diff --git a/client/internal/engine.go b/client/internal/engine.go index 0a62e7326..fc14607be 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -664,10 +664,6 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) } e.wgDevice.Store(e.wgInterface.GetWGDevice()) - // Set up notrack rules immediately after proxy is listening to prevent - // conntrack entries from being created before the rules are in place - e.setupWGProxyNoTrack() - // Start after interface is up since port may have been resolved from 0 or changed if occupied e.shutdownWg.Add(1) go func() { @@ -805,23 +801,6 @@ func (e *Engine) initFirewall() error { return nil } -// setupWGProxyNoTrack configures connection tracking exclusion for WireGuard proxy traffic. -// This prevents conntrack/MASQUERADE from affecting loopback traffic between WireGuard and the eBPF proxy. -func (e *Engine) setupWGProxyNoTrack() { - if e.firewall == nil { - return - } - - proxyPort := e.wgInterface.GetProxyPort() - if proxyPort == 0 { - return - } - - if err := e.firewall.SetupEBPFProxyNoTrack(proxyPort, uint16(e.config.WgPort)); err != nil { - log.Warnf("failed to setup ebpf proxy notrack: %v", err) - } -} - func (e *Engine) blockLanAccess() { if e.config.BlockInbound { // no need to set up extra deny rules if inbound is already blocked in general @@ -1064,7 +1043,11 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error { // back to empty if the FQDN doesn't have the expected shape. dnsName = extractDNSDomainFromFQDN(pc.GetFqdn()) } - result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName) + // With the firewall disabled there is no ACL manager to program, so + // RoutesFirewallRules would be built and then dropped. On a peer that + // routes many network resources that is the single most expensive + // step of the sync. + result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName, e.config.DisableFirewall) if err != nil { return fmt.Errorf("decode network map envelope: %w", err) } @@ -2208,10 +2191,7 @@ func (e *Engine) close() { } func (e *Engine) newWgIface() (*iface.WGIface, error) { - transportNet, err := e.newStdNet() - if err != nil { - log.Errorf("failed to create pion's stdnet: %s", err) - } + transportNet := e.newStdNet() opts := iface.WGIFaceOpts{ IFaceName: e.config.WgIfaceName, diff --git a/client/internal/engine_privileged_test.go b/client/internal/engine_privileged_test.go index 1b047e017..2db0cd5ed 100644 --- a/client/internal/engine_privileged_test.go +++ b/client/internal/engine_privileged_test.go @@ -12,12 +12,12 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/google/uuid" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "google.golang.org/grpc" "google.golang.org/grpc/keepalive" @@ -27,6 +27,7 @@ import ( "github.com/netbirdio/netbird/client/iface/wgaddr" "github.com/netbirdio/netbird/client/internal/dns" "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/profilemanager" nbssh "github.com/netbirdio/netbird/client/ssh" "github.com/netbirdio/netbird/client/system" nbdns "github.com/netbirdio/netbird/dns" @@ -81,6 +82,7 @@ func TestEngine_SSH(t *testing.T) { WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), WgPrivateKey: key, WgPort: 33100, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, ServerSSHAllowed: true, MTU: iface.DefaultMTU, SSHKey: sshKey, @@ -204,11 +206,12 @@ func TestEngine_Sync(t *testing.T) { } relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) engine := NewEngine(ctx, cancel, &EngineConfig{ - WgIfaceName: "utun103", - WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), - WgPrivateKey: key, - WgPort: 33100, - MTU: iface.DefaultMTU, + WgIfaceName: "utun103", + WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), + WgPrivateKey: key, + WgPort: 33100, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, + MTU: iface.DefaultMTU, }, EngineServices{ SignalClient: &signal.MockClient{}, MgmClient: &mgmt.MockClient{SyncFunc: syncFunc}, @@ -412,11 +415,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin wgPort := 33100 + i conf := &EngineConfig{ - WgIfaceName: ifaceName, - WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), - WgPrivateKey: key, - WgPort: wgPort, - MTU: iface.DefaultMTU, + WgIfaceName: ifaceName, + WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), + WgPrivateKey: key, + WgPort: wgPort, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, + MTU: iface.DefaultMTU, } relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) diff --git a/client/internal/engine_sessionwatch.go b/client/internal/engine_sessionwatch.go index a46d73f87..05b46a465 100644 --- a/client/internal/engine_sessionwatch.go +++ b/client/internal/engine_sessionwatch.go @@ -1,4 +1,4 @@ -//go:build !js +//go:build !js && !android package internal @@ -7,10 +7,12 @@ import ( "github.com/netbirdio/netbird/client/internal/peer" ) -// newSessionWatcher returns the real SSO session expiry watcher for every -// non-wasm build. The js/wasm build gets a no-op stub from -// engine_sessionwatch_js.go so the sessionwatch package (and its timer -// machinery) never links into the wasm binary. +// newSessionWatcher returns the real SSO session expiry watcher. The js/wasm +// build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch +// package (and its timer machinery) never links into the wasm binary; the +// android build gets a deadline-only watcher from +// engine_sessionwatch_android.go because the app schedules the warnings +// itself. func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { return sessionwatch.New(recorder) } diff --git a/client/internal/engine_sessionwatch_android.go b/client/internal/engine_sessionwatch_android.go new file mode 100644 index 000000000..8317f9165 --- /dev/null +++ b/client/internal/engine_sessionwatch_android.go @@ -0,0 +1,12 @@ +//go:build android + +package internal + +import ( + "github.com/netbirdio/netbird/client/internal/auth/sessionwatch" + "github.com/netbirdio/netbird/client/internal/peer" +) + +func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { + return sessionwatch.NewDeadlineOnly(recorder) +} diff --git a/client/internal/engine_stdnet.go b/client/internal/engine_stdnet.go index 1ebb5779c..86f6d297a 100644 --- a/client/internal/engine_stdnet.go +++ b/client/internal/engine_stdnet.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func (e *Engine) newStdNet() (*stdnet.Net, error) { +func (e *Engine) newStdNet() *stdnet.Net { return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList) } diff --git a/client/internal/engine_stdnet_android.go b/client/internal/engine_stdnet_android.go index de3c80bcf..b14deeadf 100644 --- a/client/internal/engine_stdnet_android.go +++ b/client/internal/engine_stdnet_android.go @@ -2,6 +2,6 @@ package internal import "github.com/netbirdio/netbird/client/internal/stdnet" -func (e *Engine) newStdNet() (*stdnet.Net, error) { +func (e *Engine) newStdNet() *stdnet.Net { return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList) } diff --git a/client/internal/engine_test.go b/client/internal/engine_test.go index b809b89c9..499b86e13 100644 --- a/client/internal/engine_test.go +++ b/client/internal/engine_test.go @@ -65,7 +65,6 @@ type MockWGIface struct { GetStatsFunc func() (map[string]configurer.WGStats, error) GetInterfaceGUIDStringFunc func() (string, error) GetProxyFunc func() wgproxy.Proxy - GetProxyPortFunc func() uint16 GetNetFunc func() *netstack.Net LastActivitiesFunc func() map[string]monotime.Time } @@ -162,13 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy { return m.GetProxyFunc() } -func (m *MockWGIface) GetProxyPort() uint16 { - if m.GetProxyPortFunc != nil { - return m.GetProxyPortFunc() - } - return 0 -} - func (m *MockWGIface) GetNet() *netstack.Net { return m.GetNetFunc() } @@ -696,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) { StatusRecorder: peer.NewRecorder("https://mgm"), }, MobileDependency{}) engine.ctx = ctx - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, @@ -904,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) { }, MobileDependency{}) engine.ctx = ctx - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, Address: wgaddr.MustParseWGAddress(wgAddr), @@ -1533,3 +1519,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) { }) } } + +func TestEngine_SyncResponsePersistence(t *testing.T) { + e := &Engine{} + + _, err := e.GetLatestSyncResponse() + require.Error(t, err, "persistence is disabled by default") + + e.SetSyncResponsePersistence(true) + e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}}) + + got, err := e.GetLatestSyncResponse() + require.NoError(t, err) + assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial()) + + e.SetSyncResponsePersistence(false) + _, err = e.GetLatestSyncResponse() + require.Error(t, err) +} diff --git a/client/internal/iface_common.go b/client/internal/iface_common.go index 8ffa0b102..d772a3a03 100644 --- a/client/internal/iface_common.go +++ b/client/internal/iface_common.go @@ -28,7 +28,6 @@ type wgIfaceBase interface { Up() (*udpmux.UniversalUDPMuxDefault, error) UpdateAddr(newAddr wgaddr.Address) error GetProxy() wgproxy.Proxy - GetProxyPort() uint16 UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error RemoveEndpointAddress(key string) error RemovePeer(peerKey string) error diff --git a/client/internal/ipcauth/forward_test.go b/client/internal/ipcauth/forward_test.go index d9adf05da..d80c293be 100644 --- a/client/internal/ipcauth/forward_test.go +++ b/client/internal/ipcauth/forward_test.go @@ -38,7 +38,7 @@ func asDaemon(t *testing.T, id Identity) { prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) selfIdentity, selfKnown = id, true - selfMayDelegate = !id.IsPrivileged() + selfMayDelegate = mayDelegate(id) } func TestCallerIdentity_DirectConnections(t *testing.T) { diff --git a/client/internal/ipcauth/identity.go b/client/internal/ipcauth/identity.go index d7d10f57d..255585821 100644 --- a/client/internal/ipcauth/identity.go +++ b/client/internal/ipcauth/identity.go @@ -18,7 +18,8 @@ import ( "google.golang.org/grpc/peer" ) -// Well-known Windows SIDs that identify a fully privileged principal. +// Well-known Windows SIDs. Only LocalSystem and BUILTIN\Administrators identify a +// privileged principal; the service accounts are shared by unrelated services. const ( sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE @@ -67,9 +68,9 @@ func (i Identity) IsWindows() bool { // user-to-root boundary. // // On Windows the decision comes from the caller's token rather than from -// account names or group RIDs: an elevated token, one of the service accounts -// the daemon itself may run as, or a token with BUILTIN\Administrators -// enabled. A UAC-filtered administrator has that group marked deny-only, and +// account names or group RIDs: an elevated token, the LocalSystem SID, or a +// token with BUILTIN\Administrators enabled. LocalService and NetworkService +// are not privileged by SID. A UAC-filtered administrator has that group marked deny-only, and // deny-only groups are dropped when the identity is captured, so such a // caller is correctly reported as unprivileged. Domain group memberships // (Domain Admins and friends) are deliberately not consulted: they say @@ -83,8 +84,7 @@ func (i Identity) IsPrivileged() bool { return true } - switch i.SID { - case sidLocalSystem, sidLocalService, sidNetworkService: + if i.SID == sidLocalSystem { return true } diff --git a/client/internal/ipcauth/identity_sameuser_test.go b/client/internal/ipcauth/identity_test.go similarity index 56% rename from client/internal/ipcauth/identity_sameuser_test.go rename to client/internal/ipcauth/identity_test.go index c98f583db..57be1b94e 100644 --- a/client/internal/ipcauth/identity_sameuser_test.go +++ b/client/internal/ipcauth/identity_test.go @@ -64,3 +64,58 @@ func TestIdentitySameUser(t *testing.T) { }) } } + +func TestIdentityIsPrivileged(t *testing.T) { + tests := []struct { + name string + id Identity + want bool + }{ + { + name: "Root", + id: Identity{UID: 0, GID: 0}, + want: true, + }, + { + name: "Non-root", + id: Identity{UID: 1000, GID: 1000}, + want: false, + }, + { + name: "Local system windows", + id: Identity{SID: sidLocalSystem}, + want: true, + }, + { + name: "Windows elevated", + id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Elevated: true}, + want: true, + }, + { + name: "Admin group windows", + id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Groups: []string{sidAdministrators}}, + want: true, + }, + { + name: "Regular user windows", + id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001"}, + want: false, + }, + { + name: "Network service windows", + id: Identity{SID: sidNetworkService}, + want: false, + }, + { + name: "Local service windows", + id: Identity{SID: sidLocalService}, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, tt.id.IsPrivileged()) + }) + } +} diff --git a/client/internal/ipcauth/privileged.go b/client/internal/ipcauth/privileged.go index 3c2e68432..54c66a5d2 100644 --- a/client/internal/ipcauth/privileged.go +++ b/client/internal/ipcauth/privileged.go @@ -45,7 +45,15 @@ func init() { // matching there would let a non-elevated shell of an administrator account // act as an administrator, which is the boundary the token check exists to // keep. - selfMayDelegate = !id.IsPrivileged() + selfMayDelegate = mayDelegate(id) +} + +// mayDelegate reports whether a daemon running as id may extend its authority to +// callers sharing its identity. The shared service accounts are excluded: their +// SID is held by unrelated services, so matching on it would grant them the +// daemon's authority. +func mayDelegate(id Identity) bool { + return !id.IsPrivileged() && id.SID != sidLocalService && id.SID != sidNetworkService } // IsDaemonSelf reports whether an identity is this very process. The JSON gateway diff --git a/client/internal/ipcauth/privileged_test.go b/client/internal/ipcauth/privileged_test.go index c1c7c1543..a6bbcf44b 100644 --- a/client/internal/ipcauth/privileged_test.go +++ b/client/internal/ipcauth/privileged_test.go @@ -98,7 +98,7 @@ func TestIsPrivilegedCaller_SelfRule(t *testing.T) { t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) selfIdentity, selfKnown = tt.self, tt.selfKnown - selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged() + selfMayDelegate = tt.selfKnown && mayDelegate(tt.self) if got := IsPrivilegedCaller(tt.caller); got != tt.want { t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t", @@ -132,3 +132,28 @@ func TestIsPrivilegedCaller_ThisProcess(t *testing.T) { t.Errorf("an unrelated identity %v was treated as privileged", other) } } + +// The shared service accounts are held by unrelated services, so a daemon running +// as one of them must not extend its authority to every process with that SID. +func TestMayDelegate(t *testing.T) { + tests := []struct { + name string + self Identity + want bool + }{ + {name: "unprivileged unix user", self: Identity{UID: 1000}, want: true}, + {name: "root", self: Identity{UID: 0}, want: false}, + {name: "unprivileged windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001"}, want: true}, + {name: "elevated windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, want: false}, + {name: "local system", self: Identity{SID: sidLocalSystem}, want: false}, + {name: "local service", self: Identity{SID: sidLocalService}, want: false}, + {name: "network service", self: Identity{SID: sidNetworkService}, want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := mayDelegate(tt.self); got != tt.want { + t.Errorf("mayDelegate(%+v) = %v, want %v", tt.self, got, tt.want) + } + }) + } +} diff --git a/client/internal/peer/conn.go b/client/internal/peer/conn.go index 83089606f..17823e043 100644 --- a/client/internal/peer/conn.go +++ b/client/internal/peer/conn.go @@ -135,9 +135,10 @@ type Conn struct { // used to store the remote Rosenpass key for Relayed connection in case of connection update from ice rosenpassRemoteKey []byte - wgProxyICE wgproxy.Proxy - wgProxyRelay wgproxy.Proxy - handshaker *Handshaker + wgProxyICE wgproxy.Proxy + wgProxyRelay wgproxy.Proxy + relayedConnRef *relayClient.Conn + handshaker *Handshaker guard *guard.Guard wg sync.WaitGroup @@ -560,7 +561,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { conn.mu.Lock() defer conn.mu.Unlock() - if conn.ctx.Err() != nil { + if conn.ctx.Err() != nil || rci.relayedConn.Context().Err() != nil { if err := rci.relayedConn.Close(); err != nil { conn.Log.Warnf("failed to close unnecessary relayed connection: %v", err) } @@ -575,7 +576,9 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err) return } - wgProxy.SetDisconnectListener(conn.onRelayDisconnected) + wgProxy.SetDisconnectListener(func() { + conn.onRelayDisconnected(rci.relayedConn) + }) conn.dumpState.NewLocalProxy() @@ -583,7 +586,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { if conn.isICEActive() { conn.Log.Debugf("do not switch to relay because current priority is: %s", conn.currentConnPriority.String()) - conn.setRelayedProxy(wgProxy) + conn.setRelayedProxy(wgProxy, rci.relayedConn) conn.statusRelay.SetConnected() conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, time.Now()) return @@ -614,15 +617,26 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) { conn.rosenpassRemoteKey = rci.rosenpassPubKey conn.currentConnPriority = conntype.Relay conn.statusRelay.SetConnected() - conn.setRelayedProxy(wgProxy) + conn.setRelayedProxy(wgProxy, rci.relayedConn) conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, updateTime) conn.Log.Infof("start to communicate with peer via relay") conn.doOnConnected(rci.rosenpassPubKey, rci.rosenpassAddr, updateTime) } -func (conn *Conn) onRelayDisconnected() { +// onRelayDisconnected reports the teardown of a relayed connection. relayedConn +// names the connection the signal belongs to, so a signal that arrives after +// its connection was replaced is ignored instead of tearing down its successor. +// A nil relayedConn means the caller does not track generations and the current +// connection is always torn down. +func (conn *Conn) onRelayDisconnected(relayedConn *relayClient.Conn) { conn.mu.Lock() defer conn.mu.Unlock() + + if relayedConn != nil && conn.relayedConnRef != relayedConn { + conn.Log.Debugf("ignoring relay disconnect of a superseded connection") + return + } + conn.handleRelayDisconnectedLocked() } @@ -646,6 +660,7 @@ func (conn *Conn) handleRelayDisconnectedLocked() { _ = conn.wgProxyRelay.CloseConn() conn.wgProxyRelay = nil } + conn.relayedConnRef = nil changed := conn.statusRelay.Get() != worker.StatusDisconnected if changed { @@ -813,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus { // // The result is a tri-state: // - ConnStatusConnected: all available transports are up -// - ConnStatusPartiallyConnected: relay is up but ICE is still pending/reconnecting +// - ConnStatusPartiallyConnected: one transport carries the traffic and the other does +// not: relay up with ICE down, or ICE up with the shared relay transport down // - ConnStatusDisconnected: no working transport func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { defer func() { @@ -830,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { } return evalConnStatus(connStatusInputs{ - forceRelay: IsForceRelayed(), - peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), - relayConnected: conn.statusRelay.Get() == worker.StatusConnected, - remoteSupportsICE: conn.handshaker.RemoteICESupported(), - iceWorkerCreated: iceWorkerCreated, - iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected, - iceInProgress: iceInProgress, + forceRelay: IsForceRelayed(), + peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), + relayConnected: conn.statusRelay.Get() == worker.StatusConnected, + relayTransportConnected: conn.workerRelay.IsTransportConnected(), + remoteSupportsICE: conn.handshaker.RemoteICESupported(), + iceWorkerCreated: iceWorkerCreated, + iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected, + iceInProgress: iceInProgress, }) } @@ -930,13 +947,14 @@ func (conn *Conn) logTraceConnState() { } } -func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy) { +func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy, relayedConn *relayClient.Conn) { if conn.wgProxyRelay != nil { if err := conn.wgProxyRelay.CloseConn(); err != nil { conn.Log.Warnf("failed to close deprecated wg proxy conn: %v", err) } } conn.wgProxyRelay = proxy + conn.relayedConnRef = relayedConn } // onWGHandshakeSuccess is called when the first WireGuard handshake is detected @@ -1044,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus { return boolToConnStatus(relayUsedAndUp) } - // ICE counts as "up" when the status is anything other than Disconnected, OR - // when a negotiation is currently in progress (so we don't spam offers while one is in flight). - iceUp := in.iceStatusConnecting || in.iceInProgress + // ICE counts as "running" when either connected or attempting to connect. + iceRunning := in.iceStatusConnected || in.iceInProgress // Relay side is acceptable if the peer doesn't rely on relay, or relay is connected. relayOK := !in.peerUsesRelay || in.relayConnected switch { - case iceUp && relayOK: + case iceRunning && relayOK: return guard.ConnStatusConnected case relayUsedAndUp: // Relay is up but ICE is down — partially connected. return guard.ConnStatusPartiallyConnected + case in.iceStatusConnected && !in.relayTransportConnected: + // ICE is up and the shared relay transport is down — offers cannot restore it. + return guard.ConnStatusPartiallyConnected default: return guard.ConnStatusDisconnected } diff --git a/client/internal/peer/conn_status.go b/client/internal/peer/conn_status.go index d6ad37b70..acf271534 100644 --- a/client/internal/peer/conn_status.go +++ b/client/internal/peer/conn_status.go @@ -17,13 +17,14 @@ const ( // tri-state connection classification. Extracted so the decision logic can be unit-tested // without constructing full Worker/Handshaker objects. type connStatusInputs struct { - forceRelay bool // NB_FORCE_RELAY or JS/WASM - peerUsesRelay bool // remote peer advertises relay support AND local has relay - relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) - remoteSupportsICE bool // remote peer sent ICE credentials - iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) - iceStatusConnecting bool // statusICE is anything other than Disconnected - iceInProgress bool // a negotiation is currently in flight + forceRelay bool // NB_FORCE_RELAY or JS/WASM + peerUsesRelay bool // remote peer advertises relay support AND local has relay + relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) + relayTransportConnected bool // the relay transport shared by all peers on that server is up + remoteSupportsICE bool // remote peer sent ICE credentials + iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) + iceStatusConnected bool // statusICE reports Connected + iceInProgress bool // a negotiation is currently in flight } // ConnStatus describe the status of a peer's connection diff --git a/client/internal/peer/conn_status_eval_test.go b/client/internal/peer/conn_status_eval_test.go index 66393cafe..a239196dc 100644 --- a/client/internal/peer/conn_status_eval_test.go +++ b/client/internal/peer/conn_status_eval_test.go @@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) { }, want: guard.ConnStatusDisconnected, }, + { + name: "force relay, relay up but the shared transport reports down", + in: connStatusInputs{ + forceRelay: true, + peerUsesRelay: true, + relayConnected: true, + relayTransportConnected: false, + // The ICE inputs are set so that the force-relay return is the only branch + // that can produce Connected here: without it the peer would fall through to + // relayUsedAndUp and report PartiallyConnected. + remoteSupportsICE: true, + iceWorkerCreated: true, + }, + want: guard.ConnStatusConnected, + }, { name: "force relay, peer does NOT use relay - disconnected forever", in: connStatusInputs{ @@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = true - in.iceStatusConnecting = true + in.relayTransportConnected = true + in.iceStatusConnected = true }, want: guard.ConnStatusConnected, }, { - name: "ICE connected, peer does NOT use relay", + name: "ICE connected, peer does NOT use relay, shared transport down", mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = false - in.iceStatusConnecting = true + in.relayTransportConnected = false + in.iceStatusConnected = true }, + // A peer that does not rely on relay is unaffected by the shared transport: + // relayOK is true, so the first arm matches before the transport is considered. want: guard.ConnStatusConnected, }, { name: "ICE InProgress only, peer does NOT use relay", mutator: func(in *connStatusInputs) { in.peerUsesRelay = false - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = true }, want: guard.ConnStatusConnected, @@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = true - in.iceStatusConnecting = false + in.relayTransportConnected = true + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusPartiallyConnected, @@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = false - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusDisconnected, }, { - name: "ICE up, peer uses relay but relay down -> partial (relay required, ICE ignored)", + name: "ICE connected, relay down for this peer but the shared transport is up -> disconnected", mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = false - in.iceStatusConnecting = true + in.relayTransportConnected = true + in.iceStatusConnected = true + }, + // The transport is fine, so the peer itself is unreachable over relay: it may have + // moved to another server, and only an offer carries its new relay address. + want: guard.ConnStatusDisconnected, + }, + { + name: "ICE connected, the shared relay transport is down -> partial", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = true + }, + // ICE carries the traffic and the relay transport is restored by the relay client's + // own guard, not by offers, so this must not trigger the aggressive retry. + want: guard.ConnStatusPartiallyConnected, + }, + { + name: "ICE only negotiating while the shared relay transport is down -> disconnected", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = false + in.iceInProgress = true + }, + // A negotiation in flight is not a working transport, so this peer has no path at + // all and must keep the aggressive retry. Calling it partially connected spends the + // ICE retry budget and parks the guard on the hourly ticker, and nothing wakes it + // when the negotiation then fails: onICEStateDisconnected is only reached once ICE + // has reached Connected (worker_ice.go onConnectionStateChange). + want: guard.ConnStatusDisconnected, + }, + { + name: "ICE down and the shared relay transport is down -> disconnected", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = false + in.iceInProgress = false }, - // relayOK = false (peer uses relay but it's down), iceUp = true - // first switch arm fails (relayOK false), relayUsedAndUp = false (relay down), - // falls into default: Disconnected. want: guard.ConnStatusDisconnected, }, { @@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = true // not actually used since peer doesn't rely on it - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusDisconnected, diff --git a/client/internal/peer/guard/guard.go b/client/internal/peer/guard/guard.go index 73bab2a89..15028d91c 100644 --- a/client/internal/peer/guard/guard.go +++ b/client/internal/peer/guard/guard.go @@ -14,7 +14,8 @@ type ConnStatus int const ( // ConnStatusDisconnected means neither ICE nor Relay is connected. ConnStatusDisconnected ConnStatus = iota - // ConnStatusPartiallyConnected means Relay is connected but ICE is not. + // ConnStatusPartiallyConnected means one transport is usable and the other is not: + // relay connected with ICE down, or ICE connected with the shared relay transport down. ConnStatusPartiallyConnected // ConnStatusConnected means all required connections are established. ConnStatusConnected @@ -87,8 +88,9 @@ func (g *Guard) SetICEConnDisconnected() { // - Connected: no action, the peer is fully reachable. // - Disconnected (neither ICE nor Relay): retries aggressively with exponential backoff (800ms doubling // up to timeout), never gives up. This ensures rapid recovery when the peer has no connectivity at all. -// - PartiallyConnected (Relay up, ICE not): retries up to 3 times with exponential backoff, then switches -// to one attempt per hour. This limits signaling traffic when relay already provides connectivity. +// - PartiallyConnected (one transport usable, the other not): retries up to 3 times +// with exponential backoff, then switches to one attempt per hour. This limits +// signaling traffic while the peer still has a working path. // // External events (relay/ICE disconnect, signal/relay reconnect, candidate changes) reset the retry // counter and backoff ticker, giving ICE a fresh chance after network conditions change. diff --git a/client/internal/peer/ice/agent.go b/client/internal/peer/ice/agent.go index c74b46d10..6cd8c48de 100644 --- a/client/internal/peer/ice/agent.go +++ b/client/internal/peer/ice/agent.go @@ -39,10 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c iceFailedTimeout := iceFailedTimeout() iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait() - transportNet, err := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) - if err != nil { - log.Errorf("failed to create pion's stdnet: %s", err) - } + transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) fac := logging.NewDefaultLoggerFactory() diff --git a/client/internal/peer/ice/stdnet.go b/client/internal/peer/ice/stdnet.go index 685ed0363..0c819ff66 100644 --- a/client/internal/peer/ice/stdnet.go +++ b/client/internal/peer/ice/stdnet.go @@ -8,6 +8,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) { +func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { return stdnet.NewNet(ctx, ifaceBlacklist) } diff --git a/client/internal/peer/ice/stdnet_android.go b/client/internal/peer/ice/stdnet_android.go index 5033ec1b9..2962ecf66 100644 --- a/client/internal/peer/ice/stdnet_android.go +++ b/client/internal/peer/ice/stdnet_android.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) { +func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist) } diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index d753ee43e..826bf6fe0 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -196,6 +196,7 @@ type Status struct { muxRelays sync.RWMutex peers map[string]State ipToKey map[string]string + activeRoutePeers map[route.HAUniqueID]string changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription signalState bool signalError error @@ -257,6 +258,7 @@ func NewRecorder(mgmAddress string) *Status { return &Status{ peers: make(map[string]State), ipToKey: make(map[string]string), + activeRoutePeers: make(map[route.HAUniqueID]string), changeNotify: make(map[string]map[string]*StatusChangeSubscription), eventStreams: make(map[string]chan *proto.SystemEvent), eventQueue: NewEventQueue(eventQueueSize), @@ -481,6 +483,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error { return nil } +func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) { + d.mux.Lock() + defer d.mux.Unlock() + d.activeRoutePeers[haID] = peer +} + +func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) { + d.mux.Lock() + defer d.mux.Unlock() + delete(d.activeRoutePeers, haID) +} + +func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string { + d.mux.RLock() + defer d.mux.RUnlock() + return maps.Clone(d.activeRoutePeers) +} + // CheckRoutes checks if the source and destination addresses are within the same route // and returns the resource ID of the route that contains the addresses func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) { diff --git a/client/internal/peer/status_test.go b/client/internal/peer/status_test.go index 82dff0d6f..b3f01b217 100644 --- a/client/internal/peer/status_test.go +++ b/client/internal/peer/status_test.go @@ -9,6 +9,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/route" ) func TestAddPeer(t *testing.T) { @@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) { status.MarkManagementDisconnected(err) assert.False(t, notified(ch), "redundant disconnect should not notify") } + +func TestActiveRoutePeers(t *testing.T) { + status := NewRecorder("https://mgm") + netA := route.HAUniqueID("net-a-10.0.0.0/24") + netB := route.HAUniqueID("net-b-10.0.0.0/24") + + status.AddActiveRoutePeer(netA, "peerA") + status.AddActiveRoutePeer(netB, "peerB") + + active := status.GetActiveRoutePeers() + assert.Equal(t, "peerA", active[netA]) + assert.Equal(t, "peerB", active[netB]) + + status.RemoveActiveRoutePeer(netA) + delete(active, netB) + + active = status.GetActiveRoutePeers() + _, ok := active[netA] + assert.False(t, ok) + assert.Equal(t, "peerB", active[netB]) +} diff --git a/client/internal/peer/worker_ice.go b/client/internal/peer/worker_ice.go index d17f6e693..5979e9bdc 100644 --- a/client/internal/peer/worker_ice.go +++ b/client/internal/peer/worker_ice.go @@ -121,11 +121,8 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) { } } - sessionID, err := NewICESessionID() - if err != nil { - w.log.Errorf("failed to create new session ID: %s", err) - } - w.sessionID = sessionID + // Keep the ID already advertised to the remote. Answers do not get a + // reply, so changing it here makes the next offer restart both sides. w.abandonNegotiation() } @@ -205,6 +202,9 @@ func (w *WorkerICE) Close() { w.muxAgent.Lock() defer w.muxAgent.Unlock() + if w.agent != nil || w.agentConnecting { + w.renewSessionID() + } if w.agent != nil { w.agentDialerCancel() if err := w.agent.Close(); err != nil { @@ -366,16 +366,23 @@ func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.C // Only the owner of the current session may reset its state: a stale dial // goroutine waking after a newer attempt must not clobber it. if w.agent == agent { - sessionID, err := NewICESessionID() - if err != nil { - w.log.Errorf("failed to create new session ID: %s", err) - } - w.sessionID = sessionID + w.renewSessionID() w.abandonNegotiation() } return sessionChanged } +// renewSessionID starts a new local session, so the remote treats our next offer +// or answer as a restart. Caller holds muxAgent. +func (w *WorkerICE) renewSessionID() { + sessionID, err := NewICESessionID() + if err != nil { + w.log.Errorf("failed to create new session ID: %s", err) + return + } + w.sessionID = sessionID +} + // abandonNegotiation drops all recorded ICE session state so the worker treats the // next offer as a fresh start instead of a duplicate of a dead negotiation. The // agent and agentConnecting flags must change together: leaving one stale wedges diff --git a/client/internal/peer/worker_ice_session_test.go b/client/internal/peer/worker_ice_session_test.go new file mode 100644 index 000000000..4858e0bc3 --- /dev/null +++ b/client/internal/peer/worker_ice_session_test.go @@ -0,0 +1,375 @@ +package peer + +import ( + "context" + "fmt" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + icemaker "github.com/netbirdio/netbird/client/internal/peer/ice" +) + +func TestWorkerICE_RemoteRestartPreservesAdvertisedSession(t *testing.T) { + w := newTestWorkerICE(t) + t.Cleanup(w.Close) + w.dialFunc = parkDial + advertised := w.SessionID() + remoteSession := ICESessionID("remote-first") + offer := OfferAnswer{ + IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"}, + SessionID: &remoteSession, + } + w.OnNewOffer(&offer) + require.True(t, w.InProgress(), "the first remote session must start ICE") + w.muxAgent.Lock() + firstAgent := w.agent + w.muxAgent.Unlock() + + // The same callback handles answers. A changed remote ID must not create + // an unannounced local ID that makes the remote restart on our next offer. + secondSession := ICESessionID("remote-restarted") + answer := offer + answer.SessionID = &secondSession + w.OnNewOffer(&answer) + assert.Equal(t, advertised, w.SessionID(), "following a remote restart must keep our advertised ID") + w.muxAgent.Lock() + secondAgent := w.agent + w.muxAgent.Unlock() + assert.NotSame(t, firstAgent, secondAgent, "the changed remote session must still rebuild ICE") + + w.OnNewOffer(&answer) + w.muxAgent.Lock() + defer w.muxAgent.Unlock() + assert.Same(t, secondAgent, w.agent, "a repeated answer must keep the replacement agent") +} + +func TestWorkerICE_LocalCloseChangesAdvertisedSession(t *testing.T) { + w := newTestWorkerICE(t) + dialStarted := make(chan struct{}) + dialDone := make(chan struct{}) + w.dialFunc = func(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) { + close(dialStarted) + defer close(dialDone) + <-ctx.Done() + return nil, ctx.Err() + } + session := ICESessionID("remote-session") + w.OnNewOffer(&OfferAnswer{ + IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"}, + SessionID: &session, + }) + <-dialStarted + advertised := w.SessionID() + w.Close() + assert.NotEqual(t, advertised, w.SessionID(), "a local teardown must tell the remote to restart") + closedSession := w.SessionID() + + // The abandoned dial goroutine cleans up after Close returned. + <-dialDone + assert.Never(t, func() bool { return w.SessionID() != closedSession }, 200*time.Millisecond, 10*time.Millisecond, + "the late cleanup of a closed negotiation must not restart again") + w.Close() + assert.Equal(t, closedSession, w.SessionID(), "closing an idle worker must not restart again") +} + +// parkDial stands in for the ICE dial. It never connects and returns once the +// negotiation is abandoned, so a test decides when a negotiation fails. +func parkDial(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func newTestSessionID(t *testing.T) ICESessionID { + t.Helper() + sid, err := NewICESessionID() + require.NoError(t, err) + return sid +} + +// handshakeSide is one end of a simulated signaling exchange. +type handshakeSide interface { + // message builds the offer or answer the side would send now. + message() OfferAnswer + // receive hands a remote offer or answer to the side's ICE logic. + receive(msg OfferAnswer) + // teardowns counts negotiations the side tore down to follow a remote restart. + teardowns() int + // failAgent ends the side's current negotiation as an ICE failure does. + failAgent() +} + +// workerSide drives a real WorkerICE. +type workerSide struct { + t *testing.T + w *WorkerICE + replaced int +} + +func newWorkerSide(t *testing.T) *workerSide { + t.Helper() + w := newTestWorkerICE(t) + w.dialFunc = parkDial + t.Cleanup(w.Close) + return &workerSide{t: t, w: w} +} + +func (s *workerSide) message() OfferAnswer { + sid := s.w.SessionID() + ufrag, pwd := s.w.GetLocalUserCredentials() + return OfferAnswer{IceCredentials: IceCredentials{UFrag: ufrag, Pwd: pwd}, SessionID: &sid} +} + +func (s *workerSide) receive(msg OfferAnswer) { + before := s.agent() + s.w.OnNewOffer(&msg) + if after := s.agent(); before != nil && after != before { + s.replaced++ + } +} + +func (s *workerSide) teardowns() int { return s.replaced } + +func (s *workerSide) agent() *icemaker.ThreadSafeAgent { + s.w.muxAgent.Lock() + defer s.w.muxAgent.Unlock() + return s.w.agent +} + +// failAgent runs the cleanup the dial goroutine or the Failed state callback +// performs when the current negotiation dies. +func (s *workerSide) failAgent() { + s.t.Helper() + s.w.muxAgent.Lock() + agent, cancel := s.w.agent, s.w.agentDialerCancel + s.w.muxAgent.Unlock() + require.NotNil(s.t, agent, "failing requires a running negotiation") + s.w.closeAgent(agent, cancel) +} + +// legacySide models a remote peer running a release from before this change: +// when it follows a remote restart it also picks a new session ID of its own, +// which it announces only with its next offer or answer. +type legacySide struct { + t *testing.T + sessionID ICESessionID + remoteID ICESessionID + hasAgent bool + replaced int +} + +func newLegacySide(t *testing.T) *legacySide { + return &legacySide{t: t, sessionID: newTestSessionID(t)} +} + +func (s *legacySide) message() OfferAnswer { + sid := s.sessionID + return OfferAnswer{ + IceCredentials: IceCredentials{UFrag: "legacyufrag", Pwd: "legacy-password-long-enough"}, + SessionID: &sid, + } +} + +func (s *legacySide) receive(msg OfferAnswer) { + if msg.SessionID == nil { + s.hasAgent = true + return + } + if s.hasAgent { + if *msg.SessionID == s.remoteID { + return + } + s.replaced++ + s.sessionID = newTestSessionID(s.t) + } + s.hasAgent = true + s.remoteID = *msg.SessionID +} + +func (s *legacySide) teardowns() int { return s.replaced } + +func (s *legacySide) failAgent() { + s.hasAgent = false + s.remoteID = "" + s.sessionID = newTestSessionID(s.t) +} + +// exchange runs one guard-driven round in the order Handshaker.Listen uses: the +// answerer handles the offer and answers with the session ID it holds +// afterwards, and the offerer handles the answer without replying. +func exchange(offerer, answerer handshakeSide) { + answerer.receive(offerer.message()) + offerer.receive(answerer.message()) +} + +// offerPattern decides which side's guard sends the offer in a round. +type offerPattern struct { + name string + picker func(round int, local, remote handshakeSide) (offerer, answerer handshakeSide) +} + +var offerPatterns = []offerPattern{ + { + // A routing peer whose relay is down keeps offering on its own. + name: "local peer offers", + picker: func(_ int, local, remote handshakeSide) (handshakeSide, handshakeSide) { + return local, remote + }, + }, + { + name: "both peers offer", + picker: func(round int, local, remote handshakeSide) (handshakeSide, handshakeSide) { + if round%2 == 0 { + return local, remote + } + return remote, local + }, + }, +} + +// assertSettles runs guard rounds and requires the pair to stop restarting +// each other: at most maxTeardowns in total, and none once half the rounds ran. +func assertSettles(t *testing.T, pattern offerPattern, local, remote handshakeSide, maxTeardowns int) { + t.Helper() + const rounds = 10 + + total := func() int { return local.teardowns() + remote.teardowns() } + start := total() + var halfway int + for round := range rounds { + if round == rounds/2 { + halfway = total() + } + offerer, answerer := pattern.picker(round, local, remote) + exchange(offerer, answerer) + } + + assert.LessOrEqual(t, total()-start, maxTeardowns, "the peers must not keep restarting each other") + assert.Equal(t, halfway, total(), "the negotiation must be stable in the later rounds") +} + +// establish runs the first offer and answer, so both sides negotiate. +func establish(t *testing.T, local, remote handshakeSide) { + t.Helper() + exchange(local, remote) + require.Zero(t, local.teardowns()+remote.teardowns(), "the first exchange must not restart anything") +} + +func TestICESession_SettlesAfterAgentFailure(t *testing.T) { + sides := []struct { + name string + remote func(t *testing.T) handshakeSide + }{ + {name: "current remote", remote: func(t *testing.T) handshakeSide { return newWorkerSide(t) }}, + {name: "legacy remote", remote: func(t *testing.T) handshakeSide { return newLegacySide(t) }}, + } + failures := []struct { + name string + fail func(local, remote handshakeSide) + }{ + {name: "remote agent fails", fail: func(_, remote handshakeSide) { remote.failAgent() }}, + {name: "local agent fails", fail: func(local, _ handshakeSide) { local.failAgent() }}, + {name: "both agents fail", fail: func(local, remote handshakeSide) { + local.failAgent() + remote.failAgent() + }}, + } + + for _, side := range sides { + for _, failure := range failures { + for _, pattern := range offerPatterns { + t.Run(fmt.Sprintf("%s/%s/%s", side.name, failure.name, pattern.name), func(t *testing.T) { + local := newWorkerSide(t) + remote := side.remote(t) + establish(t, local, remote) + + failure.fail(local, remote) + assertSettles(t, pattern, local, remote, 2) + }) + } + } + } +} + +// TestICESession_LocalCloseRestartsRemote covers an explicit teardown, as on a +// WireGuard handshake timeout. The remote must start over as well, or it keeps +// answering from the negotiation this side just abandoned. +func TestICESession_LocalCloseRestartsRemote(t *testing.T) { + for _, pattern := range offerPatterns { + t.Run(pattern.name, func(t *testing.T) { + local := newWorkerSide(t) + remote := newWorkerSide(t) + establish(t, local, remote) + + local.w.Close() + assertSettles(t, pattern, local, remote, 1) + assert.Equal(t, 1, remote.teardowns(), "the remote must restart its negotiation exactly once") + }) + } +} + +func TestICESession_DuplicateMessagesKeepNegotiation(t *testing.T) { + local := newWorkerSide(t) + remote := newWorkerSide(t) + + offer := local.message() + remote.receive(offer) + answer := remote.message() + local.receive(answer) + + // Signaling may deliver the same message again, and a peer answers every + // offer, including repeats of one it already handled. + remote.receive(offer) + local.receive(answer) + local.receive(remote.message()) + + assert.Zero(t, local.teardowns(), "a repeated answer must not restart the negotiation") + assert.Zero(t, remote.teardowns(), "a repeated offer must not restart the negotiation") +} + +// TestICESession_RemoteWithoutSessionIDKeepsNegotiation covers remote peers +// too old to send session IDs: once negotiating, their messages cannot tell a +// restart from a repeat, so they must not tear anything down. +func TestICESession_RemoteWithoutSessionIDKeepsNegotiation(t *testing.T) { + local := newWorkerSide(t) + unversioned := OfferAnswer{IceCredentials: IceCredentials{UFrag: "oldufrag", Pwd: "old-password-long-enough"}} + + local.receive(unversioned) + require.NotNil(t, local.agent(), "a message without a session ID must still start ICE") + advertised := local.w.SessionID() + + for range 3 { + local.receive(unversioned) + } + assert.Zero(t, local.teardowns(), "messages without a session ID must not restart the negotiation") + assert.Equal(t, advertised, local.w.SessionID(), "the advertised session must not change") +} + +// TestWorkerICE_StaleCleanupKeepsAdvertisedSession covers the cleanup of a +// replaced negotiation finishing late, from its dial goroutine or its Closed +// state callback. It must neither pick a new session ID, an unannounced local +// restart, nor disturb the negotiation that replaced it. +func TestWorkerICE_StaleCleanupKeepsAdvertisedSession(t *testing.T) { + local := newWorkerSide(t) + remote := newWorkerSide(t) + establish(t, local, remote) + + local.w.muxAgent.Lock() + oldAgent, oldCancel := local.w.agent, local.w.agentDialerCancel + local.w.muxAgent.Unlock() + + remote.failAgent() + exchange(local, remote) + require.Equal(t, 1, local.teardowns(), "the local side must follow the remote restart") + advertised := local.w.SessionID() + current := local.agent() + + local.w.closeAgent(oldAgent, oldCancel) + + assert.Equal(t, advertised, local.w.SessionID(), "a stale cleanup must not change the advertised session") + assert.Same(t, current, local.agent(), "a stale cleanup must keep the current negotiation") + assertSettles(t, offerPatterns[1], local, remote, 0) +} diff --git a/client/internal/peer/worker_relay.go b/client/internal/peer/worker_relay.go index 0402992c9..694207847 100644 --- a/client/internal/peer/worker_relay.go +++ b/client/internal/peer/worker_relay.go @@ -3,7 +3,6 @@ package peer import ( "context" "errors" - "net" "net/netip" "sync" "sync/atomic" @@ -14,7 +13,7 @@ import ( ) type RelayConnInfo struct { - relayedConn net.Conn + relayedConn *relayClient.Conn rosenpassPubKey []byte rosenpassAddr string } @@ -27,7 +26,7 @@ type WorkerRelay struct { conn *Conn relayManager *relayClient.Manager - relayedConn net.Conn + relayedConn *relayClient.Conn relayLock sync.Mutex relaySupportedOnRemotePeer atomic.Bool @@ -80,12 +79,7 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) { w.relayedConn = relayedConn w.relayLock.Unlock() - err = w.relayManager.AddCloseListener(srv, w.onRelayClientDisconnected) - if err != nil { - log.Errorf("failed to add close listener: %s", err) - _ = relayedConn.Close() - return - } + go w.watchRelayedConn(relayedConn) w.log.Debugf("peer conn opened via Relay: %s", srv) go w.conn.onRelayConnectionIsReady(RelayConnInfo{ @@ -107,14 +101,21 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool { return w.relayManager.HasRelayAddress() } +func (w *WorkerRelay) IsTransportConnected() bool { + return w.relayManager.Ready() +} + func (w *WorkerRelay) CloseConn() { w.relayLock.Lock() - defer w.relayLock.Unlock() - if w.relayedConn == nil { + conn := w.relayedConn + w.relayedConn = nil + w.relayLock.Unlock() + + if conn == nil { return } - if err := w.relayedConn.Close(); err != nil { + if err := conn.Close(); err != nil { w.log.Warnf("failed to close relay connection: %v", err) } } @@ -133,6 +134,8 @@ func (w *WorkerRelay) preferredRelayServer(myRelayAddress, remoteRelayAddress st return remoteRelayAddress } -func (w *WorkerRelay) onRelayClientDisconnected() { - go w.conn.onRelayDisconnected() +func (w *WorkerRelay) watchRelayedConn(relayedConn *relayClient.Conn) { + <-relayedConn.Context().Done() + + w.conn.onRelayDisconnected(relayedConn) } diff --git a/client/internal/profilemanager/active_state_test.go b/client/internal/profilemanager/active_state_test.go new file mode 100644 index 000000000..3b7fcd29c --- /dev/null +++ b/client/internal/profilemanager/active_state_test.go @@ -0,0 +1,76 @@ +package profilemanager + +import ( + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// Regression test: a concurrent Get and Set of the ActiveProfileState will +// fail on Windows since the write is a temp file renamed over an open file. +// Windows will refuse to replace a file another handle holds open by default. +func TestActiveProfileState_ReadsDoNotBreakAConcurrentWrite(t *testing.T) { + withTempConfigDir(t, func(configDir string) { + withPatchedGlobals(t, configDir, func() { + sm := &ServiceManager{} + require.NoError(t, sm.CreateDefaultProfile()) + require.NoError(t, sm.SetActiveProfileStateToDefault()) + + const switched = ID("0123456789abcdef0123456789abcdef") + const rounds = 50 + + var wg sync.WaitGroup + errs := make(chan error, 128) + + for i := 0; i < 8; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for r := 0; r < rounds; r++ { + state, err := sm.GetActiveProfileState() + if err != nil { + errs <- fmt.Errorf("read: %w", err) + return + } + if state.ID != defaultProfileName && state.ID != switched { + errs <- fmt.Errorf("read: active profile is %q, which no writer wrote", state.ID) + return + } + } + }() + } + + for i := 0; i < 2; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for r := 0; r < rounds; r++ { + id := switched + if r%2 == 0 { + id = defaultProfileName + } + if err := sm.SetActiveProfileState(&ActiveProfileState{ID: id, Username: "testuser"}); err != nil { + errs <- fmt.Errorf("switch: %w", err) + return + } + } + }() + } + + wg.Wait() + close(errs) + + for err := range errs { + assert.NoError(t, err, "a switch and a read of the active profile state must not collide") + } + + state, err := sm.GetActiveProfileState() + require.NoError(t, err) + assert.Contains(t, []ID{defaultProfileName, switched}, state.ID, + "the file holds whichever switch landed last, not a mix of the two") + }) + }) +} diff --git a/client/internal/relay/relay.go b/client/internal/relay/relay.go index 051717608..f0c65301e 100644 --- a/client/internal/relay/relay.go +++ b/client/internal/relay/relay.go @@ -201,11 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri } }() - net, err := stdnet.NewNet(ctx, nil) - if err != nil { - probeErr = fmt.Errorf("new net: %w", err) - return - } + net := stdnet.NewNet(ctx, nil) client, err := stun.DialURI(uri, &stun.DialConfig{ Net: net, @@ -290,11 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri } }() - net, err := stdnet.NewNet(ctx, nil) - if err != nil { - probeErr = fmt.Errorf("new net: %w", err) - return - } + net := stdnet.NewNet(ctx, nil) cfg := &turn.ClientConfig{ STUNServerAddr: turnServerAddr, TURNServerAddr: turnServerAddr, diff --git a/client/internal/routemanager/client/client.go b/client/internal/routemanager/client/client.go index c691c54f8..973cf1ab8 100644 --- a/client/internal/routemanager/client/client.go +++ b/client/internal/routemanager/client/client.go @@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err) } + w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer) if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil { log.Warnf("Failed to update peer state: %v", err) } @@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { } func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error { + w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID()) if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil { log.Warnf("Failed to update peer state: %v", err) } diff --git a/client/internal/routemanager/manager_test.go b/client/internal/routemanager/manager_test.go index 18b44820a..a1624cf46 100644 --- a/client/internal/routemanager/manager_test.go +++ b/client/internal/routemanager/manager_test.go @@ -8,6 +8,7 @@ import ( "net/netip" "testing" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" @@ -406,10 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { peerPrivateKey, _ := wgtypes.GeneratePrivateKey() - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun43%d", n), Address: wgaddr.MustParseWGAddress("100.65.65.2/24"), diff --git a/client/internal/routemanager/systemops/systemops_generic_test.go b/client/internal/routemanager/systemops/systemops_generic_test.go index c4f739c30..5b569ebd6 100644 --- a/client/internal/routemanager/systemops/systemops_generic_test.go +++ b/client/internal/routemanager/systemops/systemops_generic_test.go @@ -15,6 +15,7 @@ import ( "syscall" "testing" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -436,8 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen peerPrivateKey, err := wgtypes.GeneratePrivateKey() require.NoError(t, err) - newNet, err := stdnet.NewNet(context.Background(), nil) - require.NoError(t, err) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: interfaceName, diff --git a/client/internal/stdnet/stdnet.go b/client/internal/stdnet/stdnet.go index 381886ac6..c3a9d3d97 100644 --- a/client/internal/stdnet/stdnet.go +++ b/client/internal/stdnet/stdnet.go @@ -45,7 +45,7 @@ type Net struct { } // NewNetWithDiscover creates a new StdNet instance. -func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) (*Net, error) { +func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) *Net { if ctx == nil { ctx = context.Background() } @@ -60,20 +60,19 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover } else { n.iFaceDiscover = newMobileIFaceDiscover(iFaceDiscover) } - return n, n.UpdateInterfaces() + return n } // NewNet creates a new StdNet instance. -func NewNet(ctx context.Context, disallowList []string) (*Net, error) { +func NewNet(ctx context.Context, disallowList []string) *Net { if ctx == nil { ctx = context.Background() } - n := &Net{ + return &Net{ iFaceDiscover: pionDiscover{}, interfaceFilter: InterfaceFilter(disallowList), ctx: ctx, } - return n, n.UpdateInterfaces() } // resolveAddr performs DNS resolution with context support and timeout. @@ -122,45 +121,18 @@ func (n *Net) resolveAddr(network, address string) (netip.AddrPort, error) { return netip.AddrPortFrom(addrs[0], uint16(port)), nil } -// UpdateInterfaces updates the internal list of network interfaces -// and associated addresses filtering them by name. -// The interfaces are discovered by an external iFaceDiscover function or by a default discoverer if the external one -// wasn't specified. -func (n *Net) UpdateInterfaces() (err error) { - n.mu.Lock() - defer n.mu.Unlock() - - return n.updateInterfaces() -} - -func (n *Net) updateInterfaces() (err error) { - allIfaces, err := n.iFaceDiscover.iFaces() - if err != nil { - return err - } - - n.interfaces = n.filterInterfaces(allIfaces) - - n.lastUpdate = time.Now() - - return nil -} - // Interfaces returns a slice of interfaces which are available on the // system func (n *Net) Interfaces() ([]*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - if time.Since(n.lastUpdate) < updateInterval { - return slices.Clone(n.interfaces), nil + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err } - if err := n.updateInterfaces(); err != nil { - return nil, fmt.Errorf("update interfaces: %w", err) - } - - return slices.Clone(n.interfaces), nil + return slices.Clone(iFaces), nil } // InterfaceByIndex returns the interface specified by index. @@ -171,7 +143,13 @@ func (n *Net) Interfaces() ([]*transport.Interface, error) { func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - for _, ifc := range n.interfaces { + + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err + } + + for _, ifc := range iFaces { if ifc.Index == index { return ifc, nil } @@ -184,7 +162,13 @@ func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) { func (n *Net) InterfaceByName(name string) (*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - for _, ifc := range n.interfaces { + + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err + } + + for _, ifc := range iFaces { if ifc.Name == name { return ifc, nil } @@ -193,6 +177,31 @@ func (n *Net) InterfaceByName(name string) (*transport.Interface, error) { return nil, fmt.Errorf("%w: %s", transport.ErrInterfaceNotFound, name) } +func (n *Net) freshInterfacesLocked() ([]*transport.Interface, error) { + if time.Since(n.lastUpdate) < updateInterval { + return n.interfaces, nil + } + + if err := n.updateInterfacesLocked(); err != nil { + return nil, fmt.Errorf("update interfaces: %w", err) + } + + return n.interfaces, nil +} + +func (n *Net) updateInterfacesLocked() error { + allIFaces, err := n.iFaceDiscover.iFaces() + if err != nil { + return err + } + + n.interfaces = n.filterInterfaces(allIFaces) + + n.lastUpdate = time.Now() + + return nil +} + func (n *Net) filterInterfaces(interfaces []*transport.Interface) []*transport.Interface { if n.interfaceFilter == nil { return interfaces diff --git a/client/internal/stdnet/stdnet_test.go b/client/internal/stdnet/stdnet_test.go new file mode 100644 index 000000000..822972f39 --- /dev/null +++ b/client/internal/stdnet/stdnet_test.go @@ -0,0 +1,136 @@ +package stdnet + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/pion/transport/v3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type countingDiscover struct { + calls int + list []*transport.Interface + err error +} + +func (d *countingDiscover) iFaces() ([]*transport.Interface, error) { + d.calls++ + if d.err != nil { + return nil, d.err + } + return d.list, nil +} + +func newTestNet(t *testing.T, d iFaceDiscover) *Net { + t.Helper() + return &Net{ + iFaceDiscover: d, + ctx: context.Background(), + } +} + +func testIFace(index int, name string) *transport.Interface { + return transport.NewInterface(net.Interface{Index: index, Name: name}) +} + +func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}} + n := newTestNet(t, d) + + require.Zero(t, d.calls, "construction must not discover interfaces") + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, 1, d.calls) + + _, err = n.Interfaces() + require.NoError(t, err) + assert.Equal(t, 1, d.calls) +} + +func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) { + n := NewNet(context.Background(), nil) + require.NotNil(t, n) + assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") +} + +func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) { + n := NewNetWithDiscover(context.Background(), nil, nil) + require.NotNil(t, n) + assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") +} + +func TestNet_InterfacesRetryAfterDiscoveryFailure(t *testing.T) { + discoverErr := errors.New("discover failed") + d := &countingDiscover{err: discoverErr} + n := newTestNet(t, d) + + _, err := n.Interfaces() + require.ErrorIs(t, err, discoverErr) + + d.err = nil + d.list = []*transport.Interface{testIFace(1, "eth0")} + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, 2, d.calls) +} + +func TestNet_InterfaceByNameRefreshes(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}} + n := newTestNet(t, d) + + ifc, err := n.InterfaceByName("eth0") + require.NoError(t, err) + assert.Equal(t, "eth0", ifc.Name) + assert.Equal(t, 1, d.calls) + + _, err = n.InterfaceByName("nope") + require.ErrorIs(t, err, transport.ErrInterfaceNotFound) +} + +func TestNet_InterfaceByIndexRefreshes(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}} + n := newTestNet(t, d) + + ifc, err := n.InterfaceByIndex(3) + require.NoError(t, err) + assert.Equal(t, "eth0", ifc.Name) + assert.Equal(t, 1, d.calls) + + _, err = n.InterfaceByIndex(99) + require.ErrorIs(t, err, transport.ErrInterfaceNotFound) +} + +func TestNet_InterfaceLookupPropagatesDiscoveryError(t *testing.T) { + discoverErr := errors.New("discover failed") + n := newTestNet(t, &countingDiscover{err: discoverErr}) + + _, err := n.InterfaceByName("eth0") + require.ErrorIs(t, err, discoverErr) + + _, err = n.InterfaceByIndex(1) + require.ErrorIs(t, err, discoverErr) +} + +func TestNet_InterfacesReturnsCopy(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}} + n := newTestNet(t, d) + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + + iFaces[0] = testIFace(2, "tampered") + + iFaces, err = n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, "eth0", iFaces[0].Name) +} diff --git a/client/internal/wincmd/system32_windows.go b/client/internal/wincmd/system32_windows.go new file mode 100644 index 000000000..36aa258b5 --- /dev/null +++ b/client/internal/wincmd/system32_windows.go @@ -0,0 +1,30 @@ +// Package wincmd locates the Windows utilities the client shells out to. +package wincmd + +import ( + "path/filepath" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" +) + +// defaultSystem32Dir is where the system directory is on every supported +// install, used only when the API that reports it fails. +const defaultSystem32Dir = `C:\Windows\System32` + +// System32 returns the full path of a Windows utility under the system +// directory. +// +// PATH is deliberately not consulted. The daemon runs as LocalSystem with an +// environment of its own, so whoever can place an entry in that PATH chooses +// which binary runs with those privileges. The system directory is read from +// the API rather than from %SystemRoot% for the same reason. +func System32(command string) string { + sysDir, err := windows.GetSystemDirectory() + if err != nil { + log.Warnf("Failed to locate the Windows system directory, falling back to %s: %v", defaultSystem32Dir, err) + sysDir = defaultSystem32Dir + } + + return filepath.Join(sysDir, command+".exe") +} diff --git a/client/internal/wincmd/system32_windows_test.go b/client/internal/wincmd/system32_windows_test.go new file mode 100644 index 000000000..0d31d7ee7 --- /dev/null +++ b/client/internal/wincmd/system32_windows_test.go @@ -0,0 +1,31 @@ +package wincmd + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSystem32IgnoresPATH(t *testing.T) { + // A directory holding something that would win a PATH lookup, in front of + // everything else: the daemon runs as LocalSystem, so a PATH entry must not + // be able to decide what it executes. + planted := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(planted, "netsh.exe"), []byte("not really netsh"), 0o600)) + t.Setenv("PATH", planted+string(os.PathListSeparator)+os.Getenv("PATH")) + + got := System32("netsh") + + assert.True(t, filepath.IsAbs(got), "the path must be absolute, got %q", got) + assert.NotContains(t, got, planted, "a PATH entry must not be consulted") + assert.True(t, strings.EqualFold(filepath.Base(got), "netsh.exe"), "unexpected file name in %q", got) + + // The system directory is what Windows reports it to be, not %SystemRoot%, + // which the same caller could have set alongside PATH. + t.Setenv("SystemRoot", planted) + assert.Equal(t, got, System32("netsh"), "%SystemRoot% must not move the lookup") +} diff --git a/client/ssh/server/executor_windows.go b/client/ssh/server/executor_windows.go index 51c995ec3..9c2969d5d 100644 --- a/client/ssh/server/executor_windows.go +++ b/client/ssh/server/executor_windows.go @@ -6,7 +6,6 @@ import ( "context" "errors" "fmt" - "os" "os/exec" "os/user" "strings" @@ -506,15 +505,37 @@ func userExists(fullUsername, username, domain string) error { return nil } -// isLocalUser determines if this is a local user vs domain user +// isLocalUser reports whether domain refers to this machine rather than to a +// Windows domain. func (pd *PrivilegeDropper) isLocalUser(domain string) bool { - hostname, err := os.Hostname() - if err != nil { - hostname = "localhost" + return isLocalDomain(domain, netbiosComputerName) +} + +// isLocalDomain compares against the NetBIOS name because Windows qualifies local +// accounts with it, and it is the DNS host name truncated to 15 characters. +// An unknown name falls back to the domain path: treating it as local could +// authenticate a same named local account instead. +// https://learn.microsoft.com/en-us/windows/win32/sysinfo/computer-names +func isLocalDomain(domain string, machineName func() (string, error)) bool { + if domain == "" || domain == "." { + return true } - return domain == "" || domain == "." || - strings.EqualFold(domain, hostname) + name, err := machineName() + if err != nil { + log.Debugf("read NetBIOS computer name: %v", err) + return false + } + return strings.EqualFold(domain, name) +} + +func netbiosComputerName() (string, error) { + buf := make([]uint16, windows.MAX_COMPUTERNAME_LENGTH+1) + size := uint32(len(buf)) + if err := windows.GetComputerNameEx(windows.ComputerNamePhysicalNetBIOS, &buf[0], &size); err != nil { + return "", fmt.Errorf("GetComputerNameEx: %w", err) + } + return windows.UTF16ToString(buf[:size]), nil } // authenticateLocalUser handles authentication for local users diff --git a/client/ssh/server/executor_windows_test.go b/client/ssh/server/executor_windows_test.go new file mode 100644 index 000000000..678ca22b7 --- /dev/null +++ b/client/ssh/server/executor_windows_test.go @@ -0,0 +1,48 @@ +//go:build windows + +package server + +import ( + "errors" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sys/windows" +) + +// Past 15 characters the DNS host name and the NetBIOS name differ, and Windows +// qualifies local accounts with the NetBIOS one. +func TestIsLocalDomain(t *testing.T) { + const dnsHostname = "WINTESTMACHINE01XYZ" // 19 characters + netbios := dnsHostname[:windows.MAX_COMPUTERNAME_LENGTH] + require.NotEqual(t, strings.ToLower(dnsHostname), strings.ToLower(netbios), + "a 19 character name must not equal its 15 character truncation") + + name := func() (string, error) { return netbios, nil } + unreadable := func() (string, error) { return "", errors.New("name unavailable") } + + tests := []struct { + name string + domain string + machineName func() (string, error) + want bool + }{ + {"empty_domain", "", unreadable, true}, + {"dot_domain", ".", unreadable, true}, + {"truncated_netbios_name", netbios, name, true}, + {"netbios_name_lowercase", strings.ToLower(netbios), name, true}, + {"untruncated_dns_host_name", dnsHostname, name, false}, + {"real_domain", "CORP", name, false}, + // Must not resolve to local: that could authenticate the wrong account. + {"unreadable_machine_name", netbios, unreadable, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, isLocalDomain(tt.domain, tt.machineName), + "classification of domain %q", tt.domain) + }) + } +} diff --git a/client/ui/frontend/src/components/DropdownMenu.tsx b/client/ui/frontend/src/components/DropdownMenu.tsx index 8cedcea03..aa3ced05f 100644 --- a/client/ui/frontend/src/components/DropdownMenu.tsx +++ b/client/ui/frontend/src/components/DropdownMenu.tsx @@ -1,6 +1,6 @@ import * as DropdownMenuPrimitive from "@radix-ui/react-dropdown-menu"; import { cva } from "class-variance-authority"; -import { Check, ChevronRight, Circle } from "lucide-react"; +import { Check, ChevronRight } from "lucide-react"; import * as React from "react"; import { cn } from "@/lib/cn"; @@ -159,19 +159,23 @@ const DropdownMenuRadioItem = React.forwardRef< - + {children} + - + - {children} )); DropdownMenuRadioItem.displayName = DropdownMenuPrimitive.RadioItem.displayName; diff --git a/client/ui/frontend/src/components/LanguagePicker.tsx b/client/ui/frontend/src/components/LanguagePicker.tsx index 35ef7d5b5..d0a95906f 100644 --- a/client/ui/frontend/src/components/LanguagePicker.tsx +++ b/client/ui/frontend/src/components/LanguagePicker.tsx @@ -89,7 +89,11 @@ export function LanguagePicker() { tabIndex={0} disabled={busy || languages.length === 0} onKeyDown={handleTriggerKeyDown} - aria-label={t("settings.general.language.label")} + aria-label={ + current + ? `${t("settings.general.language.label")}: ${labelFor(current)}` + : t("settings.general.language.label") + } aria-haspopup={"listbox"} aria-expanded={open} className={cn( diff --git a/client/ui/frontend/src/components/ReadySignal.tsx b/client/ui/frontend/src/components/ReadySignal.tsx index 0d040cabc..6a98b30ea 100644 --- a/client/ui/frontend/src/components/ReadySignal.tsx +++ b/client/ui/frontend/src/components/ReadySignal.tsx @@ -1,4 +1,5 @@ import { useEffect, useRef } from "react"; +import { useSearchParams } from "react-router-dom"; import { Events } from "@wailsio/runtime"; import { useStatus } from "@/contexts/StatusContext.tsx"; @@ -6,13 +7,15 @@ const EVENT_WINDOW_PAINTED = "netbird:window-painted"; export const ReadySignal = () => { const { isReady } = useStatus(); - const sent = useRef(false); + const [params] = useSearchParams(); + const generation = params.get("gen") ?? ""; + const sent = useRef(null); useEffect(() => { - if (!isReady || sent.current) return; - sent.current = true; - void Events.Emit(EVENT_WINDOW_PAINTED); - }, [isReady]); + if (!isReady || sent.current === generation) return; + sent.current = generation; + void Events.Emit(EVENT_WINDOW_PAINTED, generation); + }, [isReady, generation]); return null; }; diff --git a/client/ui/frontend/src/components/ThemePicker.tsx b/client/ui/frontend/src/components/ThemePicker.tsx index c3dc70d1e..3acb11a99 100644 --- a/client/ui/frontend/src/components/ThemePicker.tsx +++ b/client/ui/frontend/src/components/ThemePicker.tsx @@ -1,18 +1,10 @@ import { useState } from "react"; import { useTranslation } from "react-i18next"; -import { ChevronDown, MonitorIcon, MoonIcon, SunMediumIcon, type LucideIcon } from "lucide-react"; -import { - DropdownMenu, - DropdownMenuContent, - DropdownMenuRadioGroup, - DropdownMenuRadioItem, - DropdownMenuTrigger, -} from "@/components/DropdownMenu"; +import { MonitorIcon, MoonIcon, SunMediumIcon, type LucideIcon } from "lucide-react"; +import { Select } from "@/components/inputs/Select"; import { HelpText } from "@/components/typography/HelpText"; import { Label } from "@/components/typography/Label"; import { useTheme, type ThemePreference } from "@/contexts/ThemeContext"; -import { useFocusVisible } from "@/hooks/useFocusVisible"; -import { cn } from "@/lib/cn"; import { errorDialog, formatErrorMessage } from "@/lib/errors"; const OPTIONS: { value: ThemePreference; icon: LucideIcon; labelKey: string }[] = [ @@ -25,16 +17,12 @@ export function ThemePicker() { const { t } = useTranslation(); const { theme, setTheme } = useTheme(); const [busy, setBusy] = useState(false); - const isFocusVisible = useFocusVisible(); - const current = OPTIONS.find((o) => o.value === theme) ?? OPTIONS[0]; - const CurrentIcon = current.icon; - - const select = async (value: string) => { + const select = async (value: ThemePreference) => { if (busy || value === theme) return; setBusy(true); try { - await setTheme(value as ThemePreference); + await setTheme(value); } catch (e) { await errorDialog({ Title: t("settings.error.saveTitle"), @@ -52,57 +40,17 @@ export function ThemePicker() { {t("settings.general.theme.help")}
- - - - - - void select(v)}> - {OPTIONS.map(({ value, icon: Icon, labelKey }) => ( - - - {t(labelKey)} - - ))} - - - + ({ + value, + icon, + label: t(`settings.troubleshooting.anonymize.${value}`), + }))} + onChange={setAnonymizeLevel} + ariaLabel={t("settings.troubleshooting.anonymize.label")} + />
/dev/null || \ - launchctl unload /Library/LaunchDaemons/netbird.plist 2>/dev/null || true - rm -f /Library/LaunchDaemons/netbird.plist - CMD - sudo: true + uninstall_preflight_steps do + run "/bin/launchctl", args: ["bootout", "system/netbird"], sudo: true, must_succeed: false + run "/bin/launchctl", args: ["unload", "/Library/LaunchDaemons/netbird.plist"], + sudo: true, must_succeed: false + remove "/Library/LaunchDaemons/netbird.plist", sudo: true end name "Netbird UI" diff --git a/client/ui/services/connection.go b/client/ui/services/connection.go index f78ce4c0f..f6a8eca72 100644 --- a/client/ui/services/connection.go +++ b/client/ui/services/connection.go @@ -205,19 +205,7 @@ func (s *Connection) Down(ctx context.Context) error { // window.open, so the SSO verification page can't pop inline. Honors $BROWSER // before the platform default. func (s *Connection) OpenURL(url string) error { - if browser := os.Getenv("BROWSER"); browser != "" { - return exec.Command(browser, url).Start() - } - switch runtime.GOOS { - case "windows": - return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start() - case "darwin": - return exec.Command("open", url).Start() - case "linux": - return exec.Command("xdg-open", url).Start() - default: - return fmt.Errorf("unsupported platform") - } + return openURL(url) } func (s *Connection) Logout(ctx context.Context, p LogoutParams) error { @@ -288,3 +276,19 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string, func (s *Connection) classifyDaemonError(err error) *ClientError { return s.classifier.classify(err) } + +func openURL(url string) error { + if browser := os.Getenv("BROWSER"); browser != "" { + return exec.Command(browser, url).Start() + } + switch runtime.GOOS { + case "windows": + return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start() + case "darwin": + return exec.Command("open", url).Start() + case "linux": + return exec.Command("xdg-open", url).Start() + default: + return fmt.Errorf("unsupported platform") + } +} diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index 24319dae0..af6d726a3 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -5,6 +5,7 @@ package services import ( "net/url" "strconv" + "strings" "sync" "sync/atomic" "time" @@ -26,6 +27,16 @@ type windowOp func(w *application.WebviewWindow, created bool) type windowCloser func(w *application.WebviewWindow) +// hideableWindow is the slice of application.Window the hide/restore bookkeeping needs. +// Narrow enough to fake in tests, which application.Window itself is not: it carries +// unexported methods. +type hideableWindow interface { + Show() application.Window + Hide() application.Window + IsVisible() bool + Name() string +} + // EventTriggerLogin asks the frontend's startLogin() to begin an SSO flow. const EventTriggerLogin = "trigger-login" @@ -37,7 +48,10 @@ const EventSettingsOpen = "netbird:settings:open" const EventWindowPainted = "netbird:window-painted" -const paintedFallback = 2 * time.Second +// generationParam carries the painted-report token in each dialog's start URL. +const generationParam = "gen" + +const paintedFallback = 3 * time.Second const headlessTeardownDelay = 2 * time.Second @@ -201,6 +215,12 @@ func DialogWindowOptions(name, title, url string, linuxIcon []byte) application. } } +// hiddenWindow records a window hidden by owner, the name of the popup that hid it. +type hiddenWindow struct { + win hideableWindow + owner string +} + type WindowManager struct { app *application.App mainWindow *application.WebviewWindow @@ -213,19 +233,35 @@ type WindowManager struct { installProgress *application.WebviewWindow welcome *application.WebviewWindow errorDialog *application.WebviewWindow - // hiddenForLogin holds windows hidden while the BrowserLogin popup is open, restored on close. - hiddenForLogin []application.Window - mu sync.Mutex - newMain func(startURL string) *application.WebviewWindow - creating map[string]bool - pendingOps map[string][]windowOp - pendingClose map[string]windowCloser - restoreGen uint64 - ready map[uint]bool + // hiddenWindows holds windows hidden while a popup owns the screen, each tagged with + // the popup that hid it so closing one popup cannot restore what another still hides. + hiddenWindows []hiddenWindow + hiding map[string]bool + // allWindows and raiseMain are the seams the hide/restore tests replace; both are nil + // in production, where the Wails app and the platform helper are used directly. + allWindows func() []hideableWindow + raiseMain func() + mu sync.Mutex + newMain func(startURL string) *application.WebviewWindow + creating map[string]bool + pendingOps map[string][]windowOp + pendingClose map[string]windowCloser + restoreGen map[string]uint64 + // painted gates showing a window: set by the frontend's first render, or by the + // fallback timer so a webview that never wakes up still becomes visible. + painted map[uint]bool + // mounted gates emitting to a window: set only by a real frontend report, since an + // event emitted to a frontend that has not subscribed yet is dropped, not queued. + mounted map[uint]bool showPending map[uint]bool pendingTab map[uint]string pendingEmits map[uint][]string fallbackTimers map[uint]*time.Timer + afterShow map[uint]func() + // generation maps a window name to the token stamped into its current start URL, so a + // painted report from a replaced window can be told apart from the live one's. + generation map[string]uint64 + lastGeneration uint64 headlessMain bool headlessTimer *time.Timer // recenterOnShow is set only on the minimal-WM/XEmbed path, where the WM neither centers nor @@ -243,11 +279,16 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo creating: map[string]bool{}, pendingOps: map[string][]windowOp{}, pendingClose: map[string]windowCloser{}, - ready: map[uint]bool{}, + restoreGen: map[string]uint64{}, + hiding: map[string]bool{}, + painted: map[uint]bool{}, + mounted: map[uint]bool{}, showPending: map[uint]bool{}, pendingTab: map[uint]string{}, pendingEmits: map[uint][]string{}, fallbackTimers: map[uint]*time.Timer{}, + afterShow: map[uint]func(){}, + generation: map[string]uint64{}, } s.watchPainted() s.watchTriggerLogin() @@ -307,13 +348,13 @@ func (s *WindowManager) OpenSettings(tab string) { s.withWindow(windowSettings, &s.settings, s.newSettingsWindow, func(w *application.WebviewWindow, _ bool) { s.mu.Lock() - ready := s.ready[w.ID()] - if !ready { + mounted := s.mounted[w.ID()] + if !mounted { s.pendingTab[w.ID()] = target } s.mu.Unlock() - if ready { + if mounted { s.app.Event.Emit(EventSettingsOpen, target) } s.showWhenReady(w) @@ -327,21 +368,37 @@ func (s *WindowManager) OpenBrowserLogin(uri string) { startURL = "/#/dialog/browser-login?uri=" + url.QueryEscape(uri) } s.withWindow(windowBrowserLogin, &s.browserLogin, func() *application.WebviewWindow { - return s.newBrowserLoginWindow(startURL) + return s.newBrowserLoginWindow(s.stampGeneration(windowBrowserLogin, startURL)) }, func(w *application.WebviewWindow, created bool) { - if created { - s.centerOnCursorScreen(w) - return - } - if uri != "" { - w.SetURL(startURL) + if !created && uri != "" { + w.SetURL(s.stampGeneration(windowBrowserLogin, startURL)) } s.centerOnCursorScreen(w) - w.Show() - w.Focus() + s.showThenOpenBrowser(w, uri) }) } +func (s *WindowManager) showThenOpenBrowser(w *application.WebviewWindow, uri string) { + if uri != "" { + s.mu.Lock() + s.afterShow[w.ID()] = func() { s.openBrowser(uri) } + s.mu.Unlock() + } + s.showWhenReady(w) +} + +func (s *WindowManager) openBrowser(uri string) { + if uri == "" { + return + } + go func() { + if err := openURL(uri); err != nil { + log.Errorf("open browser for SSO login: %v", err) + s.OpenError(s.title("browserLogin.openFailedTitle"), err.Error(), "") + } + }() +} + func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.WebviewWindow { s.hideOtherWindows(windowBrowserLogin) opts := DialogWindowOptions(windowBrowserLogin, s.title("window.title.signIn"), startURL, s.linuxIcon) @@ -360,12 +417,14 @@ func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.Webv if userClosed { s.browserLogin = nil } + s.forgetWindowLocked(w) s.mu.Unlock() if userClosed { - s.restoreHiddenWindows() + s.restoreHiddenWindows(windowBrowserLogin) s.app.Event.Emit(EventBrowserLoginCancel) } }) + s.armReady(w) return w } @@ -386,13 +445,11 @@ func (s *WindowManager) InstallProgressWindow() *application.WebviewWindow { } func (s *WindowManager) CloseBrowserLogin() { - // The WindowClosing hook no-ops on a programmatic close, so restore here — - // but only if a popup was actually open. The frontend calls this even when no - // popup was ever shown (e.g. resetDialog() after an early RequestExtend failure, - // or connection.ts's catch path), and hiddenForLogin is shared with - // OpenInstallProgress, so an unconditional restore could re-show windows a - // still-running install-progress is hiding. - s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoreAndClose) + // The WindowClosing hook no-ops on a programmatic close, so the closer restores. + // The frontend calls this even when no popup was ever shown (resetDialog() after an + // early RequestExtend failure, or connection.ts's catch path); closeWindow skips the + // closer then, and an owner-scoped restore cannot touch what install-progress hides. + s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin)) } // OpenSessionExpiration shows the countdown warning on the cursor's display; seconds seeds @@ -404,16 +461,13 @@ func (s *WindowManager) OpenSessionExpiration(seconds int, deadlineUnixMilli int startURL += "&deadline=" + strconv.FormatInt(deadlineUnixMilli, 10) } s.withWindow(windowSessionExpiration, &s.sessionExpiration, func() *application.WebviewWindow { - return s.newSessionExpirationWindow(startURL) + return s.newSessionExpirationWindow(s.stampGeneration(windowSessionExpiration, startURL)) }, func(w *application.WebviewWindow, created bool) { - if created { - s.centerOnCursorScreen(w) - return + if !created { + w.SetURL(s.stampGeneration(windowSessionExpiration, startURL)) } - w.SetURL(startURL) s.centerOnCursorScreen(w) - w.Show() - w.Focus() + s.showWhenReady(w) }) } @@ -427,8 +481,10 @@ func (s *WindowManager) newSessionExpirationWindow(startURL string) *application if s.sessionExpiration == w { s.sessionExpiration = nil } + s.forgetWindowLocked(w) s.mu.Unlock() }) + s.armReady(w) return w } @@ -440,20 +496,20 @@ func (s *WindowManager) CloseSessionExpiration() { // closes the browser-login popup and the session-expiration window together. func (s *WindowManager) CloseRenewFlow() { s.mu.Lock() - bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoreAndClose) + bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin)) se := s.takeWindowLocked(windowSessionExpiration, &s.sessionExpiration, closeOnly) if se != nil { - kept := s.hiddenForLogin[:0] - for _, w := range s.hiddenForLogin { - if w != se { - kept = append(kept, w) + kept := s.hiddenWindows[:0] + for _, hidden := range s.hiddenWindows { + if !sameWindow(hidden.win, se) { + kept = append(kept, hidden) } } - s.hiddenForLogin = kept + s.hiddenWindows = kept } s.mu.Unlock() - s.restoreHiddenWindows() + s.restoreHiddenWindows(windowBrowserLogin) // Close after unlock so the re-entrant handlers can take s.mu. if bl != nil { bl.Close() @@ -471,14 +527,12 @@ func (s *WindowManager) OpenInstallProgress(version string) { startURL = "/#/dialog/install-progress?version=" + url.QueryEscape(version) } s.withWindow(windowInstallProgress, &s.installProgress, func() *application.WebviewWindow { - return s.newInstallProgressWindow(startURL) + return s.newInstallProgressWindow(s.stampGeneration(windowInstallProgress, startURL)) }, func(w *application.WebviewWindow, created bool) { if !created { - w.SetURL(startURL) - w.Show() - w.Focus() + w.SetURL(s.stampGeneration(windowInstallProgress, startURL)) } - s.centerWhenReady(w) + s.showWhenReady(w) }) } @@ -489,32 +543,33 @@ func (s *WindowManager) newInstallProgressWindow(startURL string) *application.W ) w.OnWindowEvent(events.Common.WindowClosing, func(_ *application.WindowEvent) { s.mu.Lock() - if s.installProgress == w { + userClosed := s.installProgress == w + if userClosed { s.installProgress = nil } + s.forgetWindowLocked(w) s.mu.Unlock() - s.restoreHiddenWindows() + if userClosed { + s.restoreHiddenWindows(windowInstallProgress) + } }) + s.armReady(w) return w } func (s *WindowManager) CloseInstallProgress() { - s.closeWindow(windowInstallProgress, &s.installProgress, closeOnly) + s.closeWindow(windowInstallProgress, &s.installProgress, s.restoringCloser(windowInstallProgress)) } // OpenWelcome shows the first-launch onboarding window. Singleton, destroyed on close. func (s *WindowManager) OpenWelcome() { - s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, created bool) { - if !created { - w.Show() - w.Focus() - } - s.centerWhenReady(w) + s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, _ bool) { + s.showWhenReady(w) }) } func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow { - opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), "/#/dialog/welcome", s.linuxIcon) + opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), s.stampGeneration(windowWelcome, "/#/dialog/welcome"), s.linuxIcon) opts.Width = 420 opts.InitialPosition = application.WindowCentered w := s.app.Window.NewWithOptions(opts) @@ -523,8 +578,10 @@ func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow { if s.welcome == w { s.welcome = nil } + s.forgetWindowLocked(w) s.mu.Unlock() }) + s.armReady(w) return w } @@ -542,14 +599,12 @@ func (s *WindowManager) OpenError(title, message, command string) { } startURL := errorDialogURL(title, message, command) s.withWindow(windowError, &s.errorDialog, func() *application.WebviewWindow { - return s.newErrorWindow(startURL) + return s.newErrorWindow(s.stampGeneration(windowError, startURL)) }, func(w *application.WebviewWindow, created bool) { if !created { - w.SetURL(startURL) - w.Show() - w.Focus() + w.SetURL(s.stampGeneration(windowError, startURL)) } - s.centerWhenReady(w) + s.showWhenReady(w) }) } @@ -562,8 +617,10 @@ func (s *WindowManager) newErrorWindow(startURL string) *application.WebviewWind if s.errorDialog == w { s.errorDialog = nil } + s.forgetWindowLocked(w) s.mu.Unlock() }) + s.armReady(w) return w } @@ -589,14 +646,14 @@ func (s *WindowManager) ShowMainAndEmit(event string) { s.ensureMain("/", func(w *application.WebviewWindow, _ bool) { id := w.ID() s.mu.Lock() - ready := s.ready[id] - if !ready { + mounted := s.mounted[id] + if !mounted { s.pendingEmits[id] = append(s.pendingEmits[id], event) } s.mu.Unlock() s.showWhenReady(w) - if ready { + if mounted { s.app.Event.Emit(event) } }) @@ -741,31 +798,66 @@ func (s *WindowManager) releaseCreationLocked(name string) { delete(s.pendingClose, name) } -func (s *WindowManager) restoreAndClose(w *application.WebviewWindow) { - s.restoreHiddenWindows() - w.Close() +func (s *WindowManager) restoringCloser(owner string) windowCloser { + return func(w *application.WebviewWindow) { + s.restoreHiddenWindows(owner) + w.Close() + } } +// armReady starts the fallback that shows w even if its frontend never reports a first +// render. The timer starts at creation, because a hidden webview can be suspended before +// it reaches WindowRuntimeReady — the very case this fallback covers. That makes the first +// budget cover webview boot as well, so the runtime-ready hook rearms it to give the +// frontend its own full budget to mount and paint. func (s *WindowManager) armReady(w *application.WebviewWindow) { if w == nil { return } + s.armPaintedFallback(w) w.RegisterHook(events.Common.WindowRuntimeReady, func(_ *application.WindowEvent) { - timer := time.AfterFunc(paintedFallback, func() { - log.Warnf("window %q never reported a first render, showing it anyway", w.Name()) - s.markReady(w) - }) - s.mu.Lock() - s.fallbackTimers[w.ID()] = timer - s.mu.Unlock() + s.armPaintedFallback(w) }) } +func (s *WindowManager) armPaintedFallback(w *application.WebviewWindow) { + id := w.ID() + timer := time.AfterFunc(paintedFallback, func() { + s.mu.Lock() + painted := s.painted[id] + s.mu.Unlock() + if painted { + return + } + log.Warnf("window %q never reported a first render, showing it anyway", w.Name()) + s.markPainted(w) + }) + + s.mu.Lock() + if prev := s.fallbackTimers[id]; prev != nil { + prev.Stop() + } + if s.painted[id] { + timer.Stop() + delete(s.fallbackTimers, id) + } else { + s.fallbackTimers[id] = timer + } + s.mu.Unlock() +} + func (s *WindowManager) watchPainted() { s.app.Event.On(EventWindowPainted, func(e *application.CustomEvent) { - if w := s.windowByName(e.Sender); w != nil { - s.markReady(w) + w := s.windowByName(e.Sender) + if w == nil { + return } + if !s.matchesGeneration(e.Sender, paintedGeneration(e.Data)) { + log.Debugf("ignoring stale painted report for window %q", e.Sender) + return + } + s.markPainted(w) + s.markMounted(w) }) } @@ -777,7 +869,7 @@ func (s *WindowManager) watchTriggerLogin() { s.headlessTimer = nil } w := s.mainWindow - ready := w != nil && s.ready[w.ID()] + ready := w != nil && s.mounted[w.ID()] s.mu.Unlock() if ready { return @@ -788,7 +880,7 @@ func (s *WindowManager) watchTriggerLogin() { if created { s.headlessMain = true } - pending := !s.ready[w.ID()] + pending := !s.mounted[w.ID()] if pending { s.pendingEmits[w.ID()] = append(s.pendingEmits[w.ID()], EventTriggerLogin) } @@ -850,18 +942,67 @@ func (s *WindowManager) forgetWindowLocked(w *application.WebviewWindow) { timer.Stop() } delete(s.fallbackTimers, id) - delete(s.ready, id) + delete(s.painted, id) + delete(s.mounted, id) delete(s.showPending, id) delete(s.pendingTab, id) delete(s.pendingEmits, id) + delete(s.afterShow, id) - kept := s.hiddenForLogin[:0] - for _, hidden := range s.hiddenForLogin { - if hidden != application.Window(w) { + kept := s.hiddenWindows[:0] + for _, hidden := range s.hiddenWindows { + if !sameWindow(hidden.win, w) { kept = append(kept, hidden) } } - s.hiddenForLogin = kept + s.hiddenWindows = kept +} + +func (s *WindowManager) stampGeneration(name, startURL string) string { + s.mu.Lock() + defer s.mu.Unlock() + s.lastGeneration++ + s.generation[name] = s.lastGeneration + return appendGeneration(startURL, s.lastGeneration) +} + +func (s *WindowManager) matchesGeneration(name string, gen uint64) bool { + s.mu.Lock() + defer s.mu.Unlock() + want, tracked := s.generation[name] + if !tracked { + return true + } + return want == gen +} + +func (s *WindowManager) hideableWindows() []hideableWindow { + if s.allWindows != nil { + return s.allWindows() + } + all := s.app.Window.GetAll() + windows := make([]hideableWindow, 0, len(all)) + for _, w := range all { + windows = append(windows, w) + } + return windows +} + +func (s *WindowManager) isMainWindow(w hideableWindow, mainWindow *application.WebviewWindow) bool { + if s.allWindows != nil { + return w != nil && w.Name() == windowMain + } + return sameWindow(w, mainWindow) +} + +func (s *WindowManager) raiseMainWindow(mainWindow *application.WebviewWindow) { + if s.raiseMain != nil { + s.raiseMain() + return + } + if mainWindow != nil { + raiseToForeground(mainWindow) + } } func (s *WindowManager) windowByName(name string) *application.WebviewWindow { @@ -872,24 +1013,50 @@ func (s *WindowManager) windowByName(name string) *application.WebviewWindow { return s.mainWindow case windowSettings: return s.settings + case windowBrowserLogin: + return s.browserLogin + case windowSessionExpiration: + return s.sessionExpiration + case windowInstallProgress: + return s.installProgress + case windowWelcome: + return s.welcome + case windowError: + return s.errorDialog default: return nil } } -func (s *WindowManager) markReady(w *application.WebviewWindow) { +func (s *WindowManager) markPainted(w *application.WebviewWindow) { id := w.ID() s.mu.Lock() - already := s.ready[id] - s.ready[id] = true + already := s.painted[id] + s.painted[id] = true wanted := s.showPending[id] - tab, hasTab := s.pendingTab[id] - emits := s.pendingEmits[id] + delete(s.showPending, id) if timer := s.fallbackTimers[id]; timer != nil { timer.Stop() delete(s.fallbackTimers, id) } - delete(s.showPending, id) + s.mu.Unlock() + + if already || !wanted { + return + } + s.showNow(w) +} + +// markMounted records that the window's frontend is subscribed, and flushes the events +// held back for it. The fallback timer never calls this: showing a blank window is +// recoverable, emitting into a frontend that cannot hear it is not. +func (s *WindowManager) markMounted(w *application.WebviewWindow) { + id := w.ID() + s.mu.Lock() + already := s.mounted[id] + s.mounted[id] = true + tab, hasTab := s.pendingTab[id] + emits := s.pendingEmits[id] delete(s.pendingTab, id) delete(s.pendingEmits, id) s.mu.Unlock() @@ -902,10 +1069,6 @@ func (s *WindowManager) markReady(w *application.WebviewWindow) { s.app.Event.Emit(EventSettingsOpen, tab) } - if wanted { - s.showNow(w) - } - for _, event := range emits { s.app.Event.Emit(event) } @@ -918,18 +1081,19 @@ func (s *WindowManager) showWhenReady(w *application.WebviewWindow) { id := w.ID() s.mu.Lock() - ready := s.ready[id] - if !ready { + painted := s.painted[id] + if !painted { s.showPending[id] = true } s.mu.Unlock() - if ready { + if painted { s.showNow(w) } } func (s *WindowManager) showNow(w *application.WebviewWindow) { + id := w.ID() s.mu.Lock() if w == s.mainWindow { s.headlessMain = false @@ -938,10 +1102,15 @@ func (s *WindowManager) showNow(w *application.WebviewWindow) { s.headlessTimer = nil } } + after := s.afterShow[id] + delete(s.afterShow, id) s.mu.Unlock() w.Show() w.Focus() s.centerWhenReady(w) + if after != nil { + after() + } } func (s *WindowManager) ShowMainAt(url string) { @@ -1070,13 +1239,19 @@ func (s *WindowManager) retitleAll() { } } +// hideOtherWindows hides every visible window except keepName, recording them against +// keepName so only its own restore brings them back. A window already hidden by an +// earlier popup is skipped, leaving it tagged to the popup that actually hid it. The +// per-owner generation catches a restore for keepName that ran between the snapshot and +// the record, in which case the windows are re-shown rather than stranded. func (s *WindowManager) hideOtherWindows(keepName string) { s.mu.Lock() - gen := s.restoreGen + s.hiding[keepName] = true + gen := s.restoreGen[keepName] s.mu.Unlock() - var hidden []application.Window - for _, w := range s.app.Window.GetAll() { + var hidden []hideableWindow + for _, w := range s.hideableWindows() { if w == nil || w.Name() == keepName || !w.IsVisible() { continue } @@ -1088,9 +1263,11 @@ func (s *WindowManager) hideOtherWindows(keepName string) { } s.mu.Lock() - restored := s.restoreGen != gen + restored := s.restoreGen[keepName] != gen if !restored { - s.hiddenForLogin = append(s.hiddenForLogin, hidden...) + for _, w := range hidden { + s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{win: w, owner: keepName}) + } } s.mu.Unlock() if !restored { @@ -1101,33 +1278,58 @@ func (s *WindowManager) hideOtherWindows(keepName string) { } } -// restoreHiddenWindows re-shows windows hidden by hideOtherWindows. If the main -// window was among them, raiseToForeground lifts it above the SSO browser, which -// still owns the foreground — a plain Show/Focus would be demoted to a taskbar -// flash and leave it stranded behind. -func (s *WindowManager) restoreHiddenWindows() { +// restoreHiddenWindows re-shows the windows owner hid, unless another popup still covers +// them, in which case they are handed to that popup. If the main window was among them, +// raiseToForeground lifts it above the SSO browser, which still owns the foreground — a +// plain Show/Focus would be demoted to a taskbar flash and leave it stranded behind. +func (s *WindowManager) restoreHiddenWindows(owner string) { s.mu.Lock() - hidden := s.hiddenForLogin - s.hiddenForLogin = nil - s.restoreGen++ mainWindow := s.mainWindow + delete(s.hiding, owner) + var restore []hideableWindow + kept := s.hiddenWindows[:0] + for _, hidden := range s.hiddenWindows { + if hidden.owner != owner { + kept = append(kept, hidden) + continue + } + if coverer, covered := s.coveringPopupLocked(hidden.win); covered { + hidden.owner = coverer + kept = append(kept, hidden) + continue + } + if hidden.win != nil { + restore = append(restore, hidden.win) + } + } + s.hiddenWindows = kept + s.restoreGen[owner]++ s.mu.Unlock() mainRestored := false - for _, w := range hidden { - if w == nil { - continue - } + for _, w := range restore { w.Show() - if w == mainWindow { + if s.isMainWindow(w, mainWindow) { mainRestored = true } } - if mainRestored && mainWindow != nil { - raiseToForeground(mainWindow) + if mainRestored { + s.raiseMainWindow(mainWindow) } } +func (s *WindowManager) coveringPopupLocked(w hideableWindow) (string, bool) { + if w == nil { + return "", false + } + for name := range s.hiding { + if name != w.Name() { + return name, true + } + } + return "", false +} + // getScreenBasedOnCursorPosition returns the cursor's display, falling back to the // main-window screen, then nil (OS-default placement). func (s *WindowManager) getScreenBasedOnCursorPosition() *application.Screen { @@ -1169,6 +1371,48 @@ func errorDialogURL(title, message, command string) string { return startURL } +// appendGeneration adds the painted-report token to a dialog start URL, keeping any +// existing query params intact across the "/#/path?params" hash-router form. +func appendGeneration(startURL string, gen uint64) string { + sep := "?" + if strings.Contains(startURL, "?") { + sep = "&" + } + return startURL + sep + generationParam + "=" + strconv.FormatUint(gen, 10) +} + +// paintedGeneration reads the token a painted report carries back, returning 0 when the +// frontend sent none (an older bundle, or the main window, which is never stamped). +func paintedGeneration(data any) uint64 { + switch v := data.(type) { + case string: + gen, err := strconv.ParseUint(v, 10, 64) + if err != nil { + return 0 + } + return gen + case float64: + return uint64(v) + case []any: + if len(v) == 0 { + return 0 + } + return paintedGeneration(v[0]) + default: + return 0 + } +} + +// sameWindow reports whether a hidden entry refers to w, comparing through the interface +// so a nil entry never matches a live window. +func sameWindow(hidden hideableWindow, w *application.WebviewWindow) bool { + if hidden == nil || w == nil { + return false + } + other, ok := hidden.(*application.WebviewWindow) + return ok && other == w +} + // u32ptr returns a pointer to v, for the optional *uint32 Wails theme fields. func u32ptr(v uint32) *uint32 { return &v } diff --git a/client/ui/services/windowmanager_test.go b/client/ui/services/windowmanager_test.go index 13c8548ab..890fba24f 100644 --- a/client/ui/services/windowmanager_test.go +++ b/client/ui/services/windowmanager_test.go @@ -17,9 +17,65 @@ func newTestWindowManager() *WindowManager { creating: map[string]bool{}, pendingOps: map[string][]windowOp{}, pendingClose: map[string]windowCloser{}, + restoreGen: map[string]uint64{}, + hiding: map[string]bool{}, + generation: map[string]uint64{}, } } +type fakeWindow struct { + name string + visible bool + shown int + hidden int +} + +func newFakeWindow(name string) *fakeWindow { + return &fakeWindow{name: name, visible: true} +} + +func (f *fakeWindow) Show() application.Window { + f.visible = true + f.shown++ + return nil +} + +func (f *fakeWindow) Hide() application.Window { + f.visible = false + f.hidden++ + return nil +} + +func (f *fakeWindow) IsVisible() bool { return f.visible } + +func (f *fakeWindow) Name() string { return f.name } + +type fakeDesktop struct { + windows []*fakeWindow + raised int +} + +func newFakeDesktop(s *WindowManager, windows ...*fakeWindow) *fakeDesktop { + d := &fakeDesktop{windows: windows} + s.allWindows = func() []hideableWindow { + all := make([]hideableWindow, 0, len(d.windows)) + for _, w := range d.windows { + all = append(all, w) + } + return all + } + s.raiseMain = func() { d.raised++ } + return d +} + +func ownersOf(hidden []hiddenWindow) []string { + owners := make([]string, 0, len(hidden)) + for _, h := range hidden { + owners = append(owners, h.owner) + } + return owners +} + func waitDone(t *testing.T, done <-chan struct{}, msg string) { t.Helper() select { @@ -339,12 +395,245 @@ func TestCloseRenewFlowDuringBrowserLoginCreationRestoresHiddenWindows(t *testin // Seeded after the call so the deferred closer, not CloseRenewFlow's own // immediate restore, is what has to drain it. A nil entry is skipped by // restoreHiddenWindows, so no Wails window is needed. - s.hiddenForLogin = []application.Window{nil} + s.hiddenWindows = []hiddenWindow{{owner: windowBrowserLogin}} return &application.WebviewWindow{} }, func(*application.WebviewWindow, bool) {}) require.Nil(t, s.browserLogin) - require.Empty(t, s.hiddenForLogin) + require.Empty(t, s.hiddenWindows) require.Empty(t, s.creating) require.Empty(t, s.pendingClose) } + +func TestHideOtherWindowsSkipsKeepNameAndInvisible(t *testing.T) { + main := newFakeWindow(windowMain) + settings := newFakeWindow(windowSettings) + settings.visible = false + popup := newFakeWindow(windowBrowserLogin) + s := newTestWindowManager() + newFakeDesktop(s, main, settings, popup) + + s.hideOtherWindows(windowBrowserLogin) + + require.False(t, main.visible) + require.Equal(t, 1, main.hidden) + require.Equal(t, 0, settings.hidden, "an already hidden window must not be recorded") + require.Equal(t, 0, popup.hidden, "the popup itself must stay visible") + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) +} + +func TestInstallDuringLoginKeepsMainHiddenUntilLoginCloses(t *testing.T) { + main := newFakeWindow(windowMain) + login := newFakeWindow(windowBrowserLogin) + install := newFakeWindow(windowInstallProgress) + install.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, login, install) + + s.hideOtherWindows(windowBrowserLogin) + require.False(t, main.visible) + + install.visible = true + s.hideOtherWindows(windowInstallProgress) + require.False(t, login.visible, "the install popup hides the login popup") + + s.restoreHiddenWindows(windowInstallProgress) + require.True(t, login.visible, "the install popup restores the login popup it hid") + require.False(t, main.visible, "the main window stays hidden for the login popup") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) + + s.restoreHiddenWindows(windowBrowserLogin) + require.True(t, main.visible) + require.Equal(t, 1, d.raised, "restoring the main window raises it above the SSO browser") + require.Empty(t, s.hiddenWindows) +} + +func TestLoginClosingUnderInstallHandsMainToInstall(t *testing.T) { + main := newFakeWindow(windowMain) + login := newFakeWindow(windowBrowserLogin) + install := newFakeWindow(windowInstallProgress) + install.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, login, install) + + s.hideOtherWindows(windowBrowserLogin) + install.visible = true + s.hideOtherWindows(windowInstallProgress) + require.False(t, main.visible) + require.False(t, login.visible, "the install popup hides the login popup") + + // The login popup closes while the install popup is still up: the main window it + // hid must not resurface under the install popup, it is handed over instead. + s.restoreHiddenWindows(windowBrowserLogin) + require.False(t, main.visible, "the install popup still covers the main window") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowInstallProgress, windowInstallProgress}, ownersOf(s.hiddenWindows)) + + s.restoreHiddenWindows(windowInstallProgress) + require.True(t, main.visible, "the install popup restores the handed-over main window") + require.Equal(t, 1, d.raised) + require.Empty(t, s.hiddenWindows) +} + +func TestInstallClosingUnderLoginHandsMainToLogin(t *testing.T) { + main := newFakeWindow(windowMain) + install := newFakeWindow(windowInstallProgress) + login := newFakeWindow(windowBrowserLogin) + login.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, install, login) + + s.hideOtherWindows(windowInstallProgress) + login.visible = true + s.hideOtherWindows(windowBrowserLogin) + require.False(t, install.visible, "the login popup hides the install popup") + + s.restoreHiddenWindows(windowInstallProgress) + require.False(t, main.visible, "the login popup still covers the main window") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowBrowserLogin, windowBrowserLogin}, ownersOf(s.hiddenWindows)) + + s.restoreHiddenWindows(windowBrowserLogin) + require.True(t, main.visible) + require.Equal(t, 1, d.raised) + require.Empty(t, s.hiddenWindows) +} + +func TestPopupClosingReshowsTheCoveringPopupItself(t *testing.T) { + main := newFakeWindow(windowMain) + install := newFakeWindow(windowInstallProgress) + login := newFakeWindow(windowBrowserLogin) + login.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, install, login) + + s.hideOtherWindows(windowInstallProgress) + login.visible = true + s.hideOtherWindows(windowBrowserLogin) + + // The login popup hid the install popup itself; closing the login popup must bring + // the install popup back rather than hand it over to its own owner. + s.restoreHiddenWindows(windowBrowserLogin) + require.True(t, install.visible, "a popup is never handed over to itself") + require.False(t, main.visible, "the main window stays with the install popup") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows)) +} + +func TestRestoreHiddenWindowsUnknownOwnerKeepsEverything(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + d := newFakeDesktop(s, main) + s.hideOtherWindows(windowBrowserLogin) + + s.restoreHiddenWindows(windowWelcome) + + require.False(t, main.visible) + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) + require.Equal(t, 0, d.raised) +} + +func TestRestoreHiddenWindowsWithoutMainDoesNotRaise(t *testing.T) { + settings := newFakeWindow(windowSettings) + s := newTestWindowManager() + d := newFakeDesktop(s, settings) + s.hideOtherWindows(windowBrowserLogin) + + s.restoreHiddenWindows(windowBrowserLogin) + + require.True(t, settings.visible) + require.Equal(t, 0, d.raised) +} + +func TestRestoreHiddenWindowsEmptyIsNoop(t *testing.T) { + s := newTestWindowManager() + require.NotPanics(t, func() { s.restoreHiddenWindows(windowBrowserLogin) }) + require.Empty(t, s.hiddenWindows) +} + +func TestHideOtherWindowsRacingOwnRestoreReshowsWhatItHid(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + d := newFakeDesktop(s, main) + enumerate := s.allWindows + // A restore for the same owner lands between the generation snapshot and the record. + s.allWindows = func() []hideableWindow { + s.restoreHiddenWindows(windowBrowserLogin) + return enumerate() + } + + s.hideOtherWindows(windowBrowserLogin) + + require.True(t, main.visible) + require.Equal(t, 1, main.hidden) + require.Empty(t, s.hiddenWindows) + require.Equal(t, 0, d.raised) +} + +func TestHideOtherWindowsIgnoresRestoreOfAnotherOwner(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + newFakeDesktop(s, main) + enumerate := s.allWindows + s.allWindows = func() []hideableWindow { + s.restoreHiddenWindows(windowInstallProgress) + return enumerate() + } + + s.hideOtherWindows(windowBrowserLogin) + + require.False(t, main.visible) + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) +} + +func TestRestoringCloserRestoresOnlyItsOwner(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + newFakeDesktop(s, main) + s.hideOtherWindows(windowBrowserLogin) + s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{owner: windowInstallProgress}) + + s.restoringCloser(windowBrowserLogin)(&application.WebviewWindow{}) + + require.True(t, main.visible) + require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows)) +} + +func TestStampGenerationTracksLatestPerWindow(t *testing.T) { + s := newTestWindowManager() + + first := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login") + require.Equal(t, "/#/dialog/browser-login?gen=1", first) + require.True(t, s.matchesGeneration(windowBrowserLogin, 1)) + + second := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login?uri=x") + require.Equal(t, "/#/dialog/browser-login?uri=x&gen=2", second) + require.False(t, s.matchesGeneration(windowBrowserLogin, 1)) + require.True(t, s.matchesGeneration(windowBrowserLogin, 2)) +} + +func TestMatchesGenerationUntrackedWindowAccepts(t *testing.T) { + s := newTestWindowManager() + require.True(t, s.matchesGeneration(windowMain, 0)) +} + +func TestPaintedGeneration(t *testing.T) { + tests := []struct { + name string + data any + want uint64 + }{ + {"string", "7", 7}, + {"float", float64(7), 7}, + {"slice", []any{"7"}, 7}, + {"empty slice", []any{}, 0}, + {"unparsable", "abc", 0}, + {"nil", nil, 0}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.want, paintedGeneration(tc.data)) + }) + } +} diff --git a/combined/cmd/root.go b/combined/cmd/root.go index 917312e57..26d6ceedb 100644 --- a/combined/cmd/root.go +++ b/combined/cmd/root.go @@ -120,7 +120,7 @@ func execute(cmd *cobra.Command, _ []string) error { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - err = shutdownServers(ctx, servers.relaySrv, servers.healthcheck, servers.stunServer, servers.mgmtSrv, servers.metricsServer) + err = shutdownServers(ctx, servers.relaySrv, servers.healthcheck, servers.stunServer, servers.mgmtSrv, servers.signalSrv, servers.metricsServer) wg.Wait() return err } @@ -399,7 +399,7 @@ func startServers(wg *sync.WaitGroup, srv *relayServer.Server, httpHealthcheck * } } -func shutdownServers(ctx context.Context, srv *relayServer.Server, httpHealthcheck *healthcheck.Server, stunServer *stun.Server, mgmtSrv mgmtServer.Server, metricsServer *sharedMetrics.Metrics) error { +func shutdownServers(ctx context.Context, srv *relayServer.Server, httpHealthcheck *healthcheck.Server, stunServer *stun.Server, mgmtSrv mgmtServer.Server, signalSrv *signalServer.Server, metricsServer *sharedMetrics.Metrics) error { var errs error if err := httpHealthcheck.Shutdown(ctx); err != nil { @@ -425,6 +425,10 @@ func shutdownServers(ctx context.Context, srv *relayServer.Server, httpHealthche } } + if signalSrv != nil { + signalSrv.Stop() + } + if metricsServer != nil { log.Infof("shutting down metrics server") if err := metricsServer.Shutdown(ctx); err != nil { diff --git a/docs/custom-domain-validation.md b/docs/custom-domain-validation.md index a70388e4d..4b0f95ae3 100644 --- a/docs/custom-domain-validation.md +++ b/docs/custom-domain-validation.md @@ -26,3 +26,8 @@ window. Restarting management does not extend a previously assigned deadline. Registrations with existing services, including services using subdomains, are retained for operator review. Management logs their account and domain IDs so an operator can identify and resolve those dependencies before cleanup. + +Manual deletion is also refused while any service uses the domain or a subdomain, +including disabled services. Delete those services or move them to another domain +before removing the registration. A refused deletion returns HTTP 412 and leaves +the domain and its services unchanged; no deletion activity event is recorded. diff --git a/docs/testing-privileged.md b/docs/testing-privileged.md index 72e8a0f8f..939770de9 100644 --- a/docs/testing-privileged.md +++ b/docs/testing-privileged.md @@ -1,7 +1,7 @@ # Privileged tests Some tests in this repo need `root` or mutate host network state: they create -TUN/WireGuard interfaces, open netlink/raw sockets, run eBPF programs, or shell +TUN/WireGuard interfaces, open netlink/raw sockets, or shell out to `ip`/`iptables`/`nft`/`ifconfig`/`route`. Running them on a developer machine would require `sudo` and could leave stray interfaces or routes behind. @@ -44,7 +44,6 @@ A test is privileged if it does any of: - creates a real interface via `iface.NewWGIFace(...).Create()`, - opens a netlink or raw socket that hard-fails without `CAP_NET_ADMIN`, -- runs an eBPF program (`ebpf.*.Listen()`), - shells out to `ip`, `iptables`, `nft`, `ifconfig`, or `route` to change state. Add the tag to the **top** of the file, combined with any existing platform diff --git a/e2e/agentnetwork/account_delete_test.go b/e2e/agentnetwork/account_delete_test.go new file mode 100644 index 000000000..a8b52766f --- /dev/null +++ b/e2e/agentnetwork/account_delete_test.go @@ -0,0 +1,292 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "slices" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// accountDeleteModel is a made-up model id the provider enumerates and prices, +// so the chat routes to the mock upstream and is metered deterministically. +const accountDeleteModel = "e2e-account-delete-model" + +// agentNetworkConfigTables are deleted with the account, in its transaction. +var agentNetworkConfigTables = []string{ + "agent_network_settings", + "agent_network_providers", + "agent_network_policies", + "agent_network_guardrails", + "agent_network_budget_rules", +} + +// TestAccountDelete_RemovesAgentNetworkState deletes an account that has a full +// Agent Network setup and has served traffic, and checks what that leaves +// behind, end to end: +// +// - the proxy stops running the account's gateway, instead of keeping its +// mappings and provider API keys in memory until it next resyncs; +// - the configuration rows go with the account, while access logs and usage +// records stay for retention; +// - the account's consumption counters are swept once the cleanup runs; +// - the gateway domain is free for another account to claim. +// +// It runs on a dedicated server, since deleting the shared account would take +// every other test down with it. +func TestAccountDelete_RemovesAgentNetworkState(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + fresh, err := harnessStartFresh(ctx, t) + require.NoError(t, err, "start dedicated combined server") + + accounts, err := fresh.API().Accounts.List(ctx) + require.NoError(t, err, "list accounts") + require.Len(t, accounts, 1, "a fresh server has exactly the bootstrapped account") + accountID := accounts[0].Id + + cluster := harness.AgentNetworkCluster + settings, err := fresh.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{ProxyAddress: &cluster}) + require.NoError(t, err, "bootstrap agent-network endpoint") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned at bootstrap") + + env := provisionAccountDeleteEnv(t, ctx, fresh, settings.Endpoint) + chatThrough(t, ctx, env) + + // Preconditions: the proxy runs the account's gateway, and the request left + // the traffic-driven rows the rest of the test expects to outlive the delete. + requireEventually(t, ctx, 60*time.Second, "proxy should run a client for the account", func() bool { + return proxyRunsAccount(t, ctx, env.proxy, accountID) + }) + requireEventually(t, ctx, accessLogIngestWindow, "the request should leave consumption, usage and access-log rows", func() bool { + counts := accountRowCounts(t, fresh, accountID, + "agent_network_consumption", "agent_network_request_usage", "agent_network_access_log") + return counts["agent_network_consumption"] > 0 && + counts["agent_network_request_usage"] > 0 && + counts["agent_network_access_log"] > 0 + }) + + require.NoError(t, fresh.API().Accounts.Delete(ctx, accountID), "delete account") + + // The proxy is told to drop the gateway. A proxy that only learns on its + // next resync keeps serving the deleted account with its provider API keys. + if !eventually(ctx, 60*time.Second, func() bool { return proxyDroppedAccount(t, ctx, env.proxy, accountID) }) { + t.Errorf("proxy still runs a client for deleted account %s\n=== proxy logs ===\n%s", + accountID, env.proxy.Logs(context.Background())) + } + + counts := accountRowCounts(t, fresh, accountID, append(slices.Clone(agentNetworkConfigTables), + "agent_network_request_usage", "agent_network_access_log")...) + for _, table := range agentNetworkConfigTables { + assert.Zero(t, counts[table], "%s rows should be deleted with the account", table) + } + assert.NotZero(t, counts["agent_network_request_usage"], "usage records should be kept") + assert.NotZero(t, counts["agent_network_access_log"], "access logs should be left for retention") + + // The cleanup's first pass runs at startup, and whether instance setup is + // open again is only re-evaluated then. + require.NoError(t, fresh.Restart(ctx), "restart combined server") + requireEventually(t, ctx, 60*time.Second, "the cleanup should sweep the deleted account's consumption counters", func() bool { + return accountRowCounts(t, fresh, accountID, "agent_network_consumption")["agent_network_consumption"] == 0 + }) + + // A new account can claim the deleted account's gateway domain: its + // settings row no longer holds the global unique index. + _, err = fresh.Bootstrap(ctx) + require.NoError(t, err, "bootstrap a second account once the first is gone") + claimed, err := fresh.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{Endpoint: &settings.Endpoint}) + require.NoError(t, err, "a new account should be able to claim the deleted account's gateway domain") + assert.Equal(t, settings.Endpoint, claimed.Endpoint, "the new account should hold the released domain") +} + +// accountDeleteEnv is a connected gateway for one account: a proxy running the +// debug endpoint, a client peer, and the resolved endpoint. +type accountDeleteEnv struct { + endpoint string + proxyIP string + client *harness.Client + proxy *harness.Proxy +} + +// provisionAccountDeleteEnv gives the server's account one of every Agent +// Network configuration row (provider, guardrail, policy, budget rule; the +// settings row is the caller's) and brings up a proxy and a client. The policy +// and budget rule switch on usage metering, so a request records consumption. +func provisionAccountDeleteEnv(t *testing.T, ctx context.Context, srv *harness.Combined, endpoint string) accountDeleteEnv { + t.Helper() + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock vLLM upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-account-delete"}) + require.NoError(t, err, "create group") + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-account-delete-client", + Type: "reusable", + ExpiresIn: 86400, + AutoGroups: []string{grp.Id}, + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + + apiKey := "sk-account-delete-e2e" + models := []api.AgentNetworkProviderModel{{Id: accountDeleteModel, InputPer1k: 0.01, OutputPer1k: 0.02}} + prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "account-delete", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &apiKey, + Enabled: ptr(true), + Models: &models, + }) + require.NoError(t, err, "create provider") + + var gr api.AgentNetworkGuardrailRequest + gr.Name = "e2e-account-delete" + gr.Checks.ModelAllowlist.Enabled = true + gr.Checks.ModelAllowlist.Models = []string{accountDeleteModel} + guard, err := srv.CreateGuardrail(ctx, gr) + require.NoError(t, err, "create guardrail") + + limits := api.AgentNetworkPolicyLimits{ + TokenLimit: api.AgentNetworkPolicyTokenLimit{ + Enabled: true, + GroupCap: 10_000_000, + UserCap: 10_000_000, + WindowSeconds: 60, + }, + } + _, err = srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-account-delete", + Enabled: ptr(true), + SourceGroups: []string{grp.Id}, + DestinationProviderIds: []string{prov.Id}, + GuardrailIds: &[]string{guard.Id}, + Limits: &limits, + }) + require.NoError(t, err, "create policy") + + _, err = srv.CreateBudgetRule(ctx, api.AgentNetworkBudgetRuleRequest{ + Name: "e2e-account-delete", + Limits: limits, + TargetGroups: &[]string{grp.Id}, + }) + require.NoError(t, err, "create budget rule") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-account-delete-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken, map[string]string{"NB_PROXY_DEBUG_ENDPOINT": "true"}) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + proxyIP, err := cl.ResolveProxyIP(ctx, endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + return accountDeleteEnv{endpoint: endpoint, proxyIP: proxyIP, client: cl, proxy: px} +} + +// chatThrough drives one chat through the gateway, retrying to absorb +// first-call tunnel and DNS jitter. +func chatThrough(t *testing.T, ctx context.Context, env accountDeleteEnv) { + t.Helper() + var code int + var body string + ok := eventually(ctx, 90*time.Second, func() bool { + c, b, err := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, + accountDeleteModel, "Reply with exactly: pong", "e2e-session-account-delete") + code, body = c, b + return err == nil && c == 200 + }) + require.True(t, ok, "chat must return 200, last got %d: %s\n=== proxy logs ===\n%s", + code, body, env.proxy.Logs(context.Background())) +} + +// proxyRunsAccount reports whether a lookup succeeded and shows the proxy +// running a client for the account. +func proxyRunsAccount(t *testing.T, ctx context.Context, px *harness.Proxy, accountID string) bool { + t.Helper() + runs, ok := lookupProxyAccount(t, ctx, px, accountID) + return ok && runs +} + +// proxyDroppedAccount reports whether a lookup succeeded and shows the proxy +// no longer running a client for the account. A failed lookup confirms +// nothing, so it keeps the caller polling rather than passing the check. +func proxyDroppedAccount(t *testing.T, ctx context.Context, px *harness.Proxy, accountID string) bool { + t.Helper() + runs, ok := lookupProxyAccount(t, ctx, px, accountID) + return ok && !runs +} + +// lookupProxyAccount asks the proxy whether it runs a client for the account; +// ok is false when the lookup itself failed. +func lookupProxyAccount(t *testing.T, ctx context.Context, px *harness.Proxy, accountID string) (runs, ok bool) { + t.Helper() + clients, err := px.DebugClients(ctx) + if err != nil { + t.Logf("proxy debug clients: %v", err) + return false, false + } + return slices.ContainsFunc(clients, func(c harness.ProxyDebugClient) bool { return c.AccountID == accountID }), true +} + +// accountRowCounts counts the account's rows in each table, read from a +// snapshot of the management store. +func accountRowCounts(t *testing.T, srv *harness.Combined, accountID string, tables ...string) map[string]int64 { + t.Helper() + dbPath, err := srv.SnapshotStoreDB(t.TempDir()) + require.NoError(t, err, "snapshot management sqlite store") + db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{}) + require.NoError(t, err, "open store snapshot") + sqlDB, err := db.DB() + require.NoError(t, err) + defer func() { _ = sqlDB.Close() }() + + counts := make(map[string]int64, len(tables)) + for _, table := range tables { + var n int64 + require.NoError(t, db.Table(table).Where("account_id = ?", accountID).Count(&n).Error, "count %s rows", table) + counts[table] = n + } + return counts +} + +// eventually polls cond every two seconds until it holds or timeout passes. +func eventually(ctx context.Context, timeout time.Duration, cond func() bool) bool { + deadline := time.Now().Add(timeout) + for { + if cond() { + return true + } + if time.Now().After(deadline) || !waitBeforeRetry(ctx, 2*time.Second) { + return false + } + } +} + +// requireEventually fails the test now if cond does not hold within timeout. +func requireEventually(t *testing.T, ctx context.Context, timeout time.Duration, msg string, cond func() bool) { + t.Helper() + require.True(t, eventually(ctx, timeout, cond), msg) +} diff --git a/e2e/harness/agentnetwork.go b/e2e/harness/agentnetwork.go index e51f2dd7a..8f688836d 100644 --- a/e2e/harness/agentnetwork.go +++ b/e2e/harness/agentnetwork.go @@ -135,6 +135,11 @@ func (c *Combined) DeleteGuardrail(ctx context.Context, id string) error { return anDelete(ctx, c, "/api/agent-network/guardrails/"+id) } +// CreateBudgetRule creates an account-level agent-network budget rule. +func (c *Combined) CreateBudgetRule(ctx context.Context, req api.AgentNetworkBudgetRuleRequest) (api.AgentNetworkBudgetRule, error) { + return anRequest[api.AgentNetworkBudgetRule](ctx, c, http.MethodPost, "/api/agent-network/budget-rules", req) +} + // CreateSettings bootstraps the account's agent-network settings row, // assigning the immutable endpoint. Exactly one of req.ProxyAddress (labeled // endpoint beneath that cluster) and req.Endpoint (self-addressed dedicated diff --git a/e2e/harness/combined.go b/e2e/harness/combined.go index e03f9f256..ea451e3ad 100644 --- a/e2e/harness/combined.go +++ b/e2e/harness/combined.go @@ -305,6 +305,33 @@ func (c *Combined) SnapshotStoreDB(dstDir string) (string, error) { return dst, nil } +// Restart stops and starts the combined container, keeping its bind-mounted +// data dir, and waits for the API again. The host port can change across a +// restart, so BaseURL and the authenticated client are refreshed. Work that +// management only does at startup (such as the agent-network cleanup's first +// pass, or re-evaluating whether instance setup is required) runs again. +func (c *Combined) Restart(ctx context.Context) error { + if err := c.container.Stop(ctx, nil); err != nil { + return fmt.Errorf("stop combined container: %w", err) + } + if err := c.container.Start(ctx); err != nil { + return fmt.Errorf("start combined container: %w", err) + } + host, err := c.container.Host(ctx) + if err != nil { + return fmt.Errorf("container host: %w", err) + } + mapped, err := c.container.MappedPort(ctx, nat.Port(combinedHTTPPort)) + if err != nil { + return fmt.Errorf("mapped port: %w", err) + } + c.BaseURL = fmt.Sprintf("http://%s:%s", host, mapped.Port()) + if c.PAT != "" { + c.api = rest.New(c.BaseURL, c.PAT) + } + return nil +} + // Logs returns the combined server container logs, for diagnostics. func (c *Combined) Logs(ctx context.Context) string { return containerLogs(ctx, c.container) diff --git a/e2e/harness/proxy.go b/e2e/harness/proxy.go index 3d709b439..ee458908b 100644 --- a/e2e/harness/proxy.go +++ b/e2e/harness/proxy.go @@ -3,13 +3,17 @@ package harness import ( + "bytes" "context" + "encoding/json" "fmt" + "io" "os" "time" "github.com/docker/docker/api/types/container" "github.com/testcontainers/testcontainers-go" + tcexec "github.com/testcontainers/testcontainers-go/exec" "github.com/testcontainers/testcontainers-go/wait" ) @@ -114,6 +118,41 @@ func StartProxy(ctx context.Context, c *Combined, proxyToken string, envOverride return &Proxy{container: ctr, workDir: workDir}, nil } +// ProxyDebugClient is one per-account embedded client the proxy runs, as the +// proxy's debug endpoint reports it. +type ProxyDebugClient struct { + AccountID string `json:"account_id"` + ServiceCount int `json:"service_count"` + ServiceKeys []string `json:"service_keys"` +} + +// DebugClients lists the per-account clients the proxy is running, through +// the proxy's own debug CLI inside the container. The proxy must be started +// with NB_PROXY_DEBUG_ENDPOINT=true. +func (p *Proxy) DebugClients(ctx context.Context) ([]ProxyDebugClient, error) { + code, reader, err := p.container.Exec(ctx, + []string{"/usr/bin/netbird-proxy", "debug", "clients", "--json"}, tcexec.Multiplexed()) + if err != nil { + return nil, fmt.Errorf("exec debug clients: %w", err) + } + out, _ := io.ReadAll(reader) + if code != 0 { + return nil, fmt.Errorf("debug clients exited %d: %s", code, string(out)) + } + // stderr is multiplexed in; the JSON document starts at the first brace. + start := bytes.IndexByte(out, '{') + if start < 0 { + return nil, fmt.Errorf("no JSON in debug clients output: %s", string(out)) + } + var resp struct { + Clients []ProxyDebugClient `json:"clients"` + } + if err := json.NewDecoder(bytes.NewReader(out[start:])).Decode(&resp); err != nil { + return nil, fmt.Errorf("decode debug clients output: %w", err) + } + return resp.Clients, nil +} + // Logs returns the proxy container logs, for diagnostics on failure. func (p *Proxy) Logs(ctx context.Context) string { return containerLogs(ctx, p.container) diff --git a/encryption/sharedkey.go b/encryption/sharedkey.go new file mode 100644 index 000000000..1509632c4 --- /dev/null +++ b/encryption/sharedkey.go @@ -0,0 +1,147 @@ +package encryption + +import ( + "fmt" + "sync" + + pb "github.com/golang/protobuf/proto" //nolint + "golang.org/x/crypto/nacl/box" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// SharedKeyCache encrypts and decrypts messages for one local private key, deriving +// the box shared key once per remote public key instead of once per message. +// +// The shared key is a pure function of the two keys, so a cached entry never goes +// stale: a different remote key is a different entry, and a different local key +// needs a different cache. Entries are only dropped to stay under maxSharedKeys. +// Every message still uses its own random nonce. +// +// The cached values are secret key material, as sensitive as the private key. +type SharedKeyCache struct { + privateKey wgtypes.Key + limit int + + mu sync.RWMutex + keys map[wgtypes.Key]*[32]byte + closed bool +} + +// NewSharedKeyCache returns a cache for messages sent and received with privateKey. +func NewSharedKeyCache(privateKey wgtypes.Key) *SharedKeyCache { + return &SharedKeyCache{ + privateKey: privateKey, + limit: maxSharedKeys, + keys: make(map[wgtypes.Key]*[32]byte), + } +} + +// Encrypt encrypts msg for peerPublicKey. It is safe for concurrent use. +func (c *SharedKeyCache) Encrypt(msg []byte, peerPublicKey wgtypes.Key) ([]byte, error) { + nonce, err := genNonce() + if err != nil { + return nil, err + } + return box.SealAfterPrecomputation(nonce[:], msg, nonce, c.sharedKey(peerPublicKey)), nil +} + +// Decrypt decrypts a message that peerPublicKey encrypted for this cache's private +// key. It is safe for concurrent use. +func (c *SharedKeyCache) Decrypt(encryptedMsg []byte, peerPublicKey wgtypes.Key) ([]byte, error) { + if len(encryptedMsg) < nonceSize { + return nil, fmt.Errorf("invalid encrypted message length") + } + + var nonce [nonceSize]byte + copy(nonce[:], encryptedMsg[:nonceSize]) + + shared, cached := c.cached(peerPublicKey) + if !cached { + shared = c.derive(peerPublicKey) + } + + opened, ok := box.OpenAfterPrecomputation(nil, encryptedMsg[nonceSize:], &nonce, shared) + if !ok { + return nil, fmt.Errorf("failed to decrypt message from peer %s", peerPublicKey.String()) + } + + // The sender key of an incoming message is not authenticated until it opens, so + // only a key that produced a valid message is cached. Forged senders cannot fill + // the cache or evict real peers. + if !cached { + c.store(peerPublicKey, shared) + } + return opened, nil +} + +// EncryptMessage marshals message and encrypts it for peerPublicKey. +func (c *SharedKeyCache) EncryptMessage(peerPublicKey wgtypes.Key, message pb.Message) ([]byte, error) { + body, err := pb.Marshal(message) + if err != nil { + return nil, fmt.Errorf("marshal message: %w", err) + } + return c.Encrypt(body, peerPublicKey) +} + +// DecryptMessage decrypts a message from peerPublicKey and unmarshals it into message. +func (c *SharedKeyCache) DecryptMessage(peerPublicKey wgtypes.Key, encryptedMessage []byte, message pb.Message) error { + body, err := c.Decrypt(encryptedMessage, peerPublicKey) + if err != nil { + return err + } + if err := pb.Unmarshal(body, message); err != nil { + return fmt.Errorf("unmarshal message from peer %s: %w", peerPublicKey.String(), err) + } + return nil +} + +// Close drops every cached shared key and stops caching new ones. Encrypt and +// Decrypt keep working afterwards by deriving the key for each message. +func (c *SharedKeyCache) Close() { + c.mu.Lock() + defer c.mu.Unlock() + c.closed = true + clear(c.keys) +} + +func (c *SharedKeyCache) sharedKey(peerPublicKey wgtypes.Key) *[32]byte { + if shared, ok := c.cached(peerPublicKey); ok { + return shared + } + + shared := c.derive(peerPublicKey) + c.store(peerPublicKey, shared) + return shared +} + +func (c *SharedKeyCache) cached(peerPublicKey wgtypes.Key) (*[32]byte, bool) { + c.mu.RLock() + defer c.mu.RUnlock() + shared, ok := c.keys[peerPublicKey] + return shared, ok +} + +// derive computes the shared key outside the lock: two goroutines racing on a new +// peer compute the same value, and holding the lock would serialise the x25519 work +// this cache avoids. +func (c *SharedKeyCache) derive(peerPublicKey wgtypes.Key) *[32]byte { + shared := new([32]byte) + box.Precompute(shared, toByte32(peerPublicKey), toByte32(c.privateKey)) + return shared +} + +func (c *SharedKeyCache) store(peerPublicKey wgtypes.Key, shared *[32]byte) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed { + return + } + if len(c.keys) >= c.limit { + // Map iteration order is random, so this evicts an arbitrary entry. + for k := range c.keys { + delete(c.keys, k) + break + } + } + c.keys[peerPublicKey] = shared +} diff --git a/encryption/sharedkey_limit.go b/encryption/sharedkey_limit.go new file mode 100644 index 000000000..b141492f6 --- /dev/null +++ b/encryption/sharedkey_limit.go @@ -0,0 +1,8 @@ +//go:build !ios && !android + +package encryption + +// maxSharedKeys bounds the cache so peers that come and go (ephemeral peers get a +// new key on every registration) cannot grow it for the lifetime of the process. +// An entry costs about 130 bytes, so a full cache is around 8 MB. +const maxSharedKeys = 1 << 16 diff --git a/encryption/sharedkey_limit_mobile.go b/encryption/sharedkey_limit_mobile.go new file mode 100644 index 000000000..f36181f94 --- /dev/null +++ b/encryption/sharedkey_limit_mobile.go @@ -0,0 +1,8 @@ +//go:build ios || android + +package encryption + +// maxSharedKeys is small on mobile, where the process runs under a tight memory +// limit. A miss only costs a fresh key derivation. An entry costs about 130 bytes, +// so a full cache is around 130 KB. +const maxSharedKeys = 1 << 10 diff --git a/encryption/sharedkey_test.go b/encryption/sharedkey_test.go new file mode 100644 index 000000000..003ccb627 --- /dev/null +++ b/encryption/sharedkey_test.go @@ -0,0 +1,184 @@ +package encryption + +import ( + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +func newKeyPair(t testing.TB) (wgtypes.Key, wgtypes.Key) { + t.Helper() + priv, err := wgtypes.GeneratePrivateKey() + require.NoError(t, err) + return priv, priv.PublicKey() +} + +// The cache must stay wire compatible with peers that use the uncached functions, +// in both directions. +func TestSharedKeyCache_InteropWithUncached(t *testing.T) { + alicePriv, alicePub := newKeyPair(t) + bobPriv, bobPub := newKeyPair(t) + alice := NewSharedKeyCache(alicePriv) + msg := []byte("offer") + + enc, err := alice.Encrypt(msg, bobPub) + require.NoError(t, err) + dec, err := Decrypt(enc, alicePub, bobPriv) + require.NoError(t, err) + assert.Equal(t, msg, dec, "uncached peer must read a cached sender's message") + + enc, err = Encrypt(msg, alicePub, bobPriv) + require.NoError(t, err) + dec, err = alice.Decrypt(enc, bobPub) + require.NoError(t, err) + assert.Equal(t, msg, dec, "cached peer must read an uncached sender's message") +} + +// Two messages to the same peer share the derived key but never the nonce, so the +// ciphertexts differ. +func TestSharedKeyCache_FreshNoncePerMessage(t *testing.T) { + priv, _ := newKeyPair(t) + _, peerPub := newKeyPair(t) + c := NewSharedKeyCache(priv) + + a, err := c.Encrypt([]byte("same"), peerPub) + require.NoError(t, err) + b, err := c.Encrypt([]byte("same"), peerPub) + require.NoError(t, err) + assert.NotEqual(t, a, b, "ciphertexts of identical plaintext must differ") + assert.Len(t, c.keys, 1, "the shared key must be derived once per peer") +} + +// A message from one peer must not decrypt under another peer's cached key. +func TestSharedKeyCache_DoesNotMixPeers(t *testing.T) { + alicePriv, alicePub := newKeyPair(t) + bobPriv, _ := newKeyPair(t) + _, carolPub := newKeyPair(t) + alice := NewSharedKeyCache(alicePriv) + + enc, err := Encrypt([]byte("hi"), alicePub, bobPriv) + require.NoError(t, err) + _, err = alice.Decrypt(enc, carolPub) + assert.Error(t, err, "a message from Bob must not open with Carol's key") +} + +func TestSharedKeyCache_RejectsShortMessage(t *testing.T) { + priv, _ := newKeyPair(t) + _, peerPub := newKeyPair(t) + _, err := NewSharedKeyCache(priv).Decrypt(make([]byte, nonceSize-1), peerPub) + assert.Error(t, err) +} + +func TestSharedKeyCache_StaysBounded(t *testing.T) { + priv, _ := newKeyPair(t) + c := NewSharedKeyCache(priv) + c.limit = 4 + + for i := 0; i < 20; i++ { + _, peerPub := newKeyPair(t) + _, err := c.Encrypt([]byte("x"), peerPub) + require.NoError(t, err) + assert.LessOrEqual(t, len(c.keys), c.limit, "cache must not grow past its cap") + } + assert.Len(t, c.keys, c.limit, "a full cache keeps evicting one entry per new peer") + c.Close() + assert.Empty(t, c.keys, "Close must drop every entry") +} + +func TestSharedKeyCache_Concurrent(t *testing.T) { + alicePriv, alicePub := newKeyPair(t) + bobPriv, bobPub := newKeyPair(t) + alice := NewSharedKeyCache(alicePriv) + bob := NewSharedKeyCache(bobPriv) + + var wg sync.WaitGroup + for i := 0; i < 16; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 50; j++ { + enc, err := alice.Encrypt([]byte("m"), bobPub) + if !assert.NoError(t, err) { + return + } + dec, err := bob.Decrypt(enc, alicePub) + if !assert.NoError(t, err) || !assert.Equal(t, []byte("m"), dec) { + return + } + } + }() + } + wg.Wait() +} + +func BenchmarkEncryptDecryptUncached(b *testing.B) { + alicePriv, alicePub := newKeyPair(b) + bobPriv, bobPub := newKeyPair(b) + msg := make([]byte, 512) + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + enc, err := Encrypt(msg, bobPub, alicePriv) + require.NoError(b, err) + _, err = Decrypt(enc, alicePub, bobPriv) + require.NoError(b, err) + } +} + +func BenchmarkEncryptDecryptCached(b *testing.B) { + alicePriv, alicePub := newKeyPair(b) + bobPriv, bobPub := newKeyPair(b) + alice := NewSharedKeyCache(alicePriv) + bob := NewSharedKeyCache(bobPriv) + msg := make([]byte, 512) + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + enc, err := alice.Encrypt(msg, bobPub) + require.NoError(b, err) + _, err = bob.Decrypt(enc, alicePub) + require.NoError(b, err) + } +} + +// A forged sender key must not populate the cache: the key of an incoming message +// is only trusted once the message opens. +func TestSharedKeyCache_FailedDecryptDoesNotCache(t *testing.T) { + alicePriv, alicePub := newKeyPair(t) + bobPriv, bobPub := newKeyPair(t) + _, forgedPub := newKeyPair(t) + alice := NewSharedKeyCache(alicePriv) + + enc, err := Encrypt([]byte("hi"), alicePub, bobPriv) + require.NoError(t, err) + + _, err = alice.Decrypt(enc, forgedPub) + require.Error(t, err) + assert.Empty(t, alice.keys, "a message that fails to open must not add a cache entry") + + _, err = alice.Decrypt(enc, bobPub) + require.NoError(t, err) + assert.Len(t, alice.keys, 1, "a message that opens caches its sender's key") +} + +// After Close the cache still works but no longer keeps key material. +func TestSharedKeyCache_ClosedDoesNotRepopulate(t *testing.T) { + alicePriv, alicePub := newKeyPair(t) + bobPriv, bobPub := newKeyPair(t) + alice := NewSharedKeyCache(alicePriv) + + _, err := alice.Encrypt([]byte("x"), bobPub) + require.NoError(t, err) + alice.Close() + assert.Empty(t, alice.keys) + + enc, err := alice.Encrypt([]byte("y"), bobPub) + require.NoError(t, err) + dec, err := Decrypt(enc, alicePub, bobPriv) + require.NoError(t, err) + assert.Equal(t, []byte("y"), dec, "a closed cache must still encrypt correctly") + assert.Empty(t, alice.keys, "a closed cache must not cache new keys") +} diff --git a/go.mod b/go.mod index e15023b49..52deafdd2 100644 --- a/go.mod +++ b/go.mod @@ -40,8 +40,8 @@ require ( github.com/aws/aws-sdk-go-v2/credentials v1.20.4 github.com/aws/aws-sdk-go-v2/service/s3 v1.87.3 github.com/c-robinson/iplib v1.0.3 + github.com/caarlos0/env/v11 v11.4.1 github.com/caddyserver/certmagic v0.21.3 - github.com/cilium/ebpf v0.19.0 github.com/coder/websocket v1.8.14 github.com/coreos/go-iptables v0.7.0 github.com/coreos/go-oidc/v3 v3.18.0 @@ -69,6 +69,7 @@ require ( github.com/google/gopacket v1.1.19 github.com/google/nftables v0.3.0 github.com/gopacket/gopacket v1.4.0 + github.com/grafana/pyroscope-go v1.4.2 github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357 github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 github.com/hashicorp/go-multierror v1.1.1 @@ -190,6 +191,7 @@ require ( github.com/caddyserver/zerossl v0.1.3 // indirect github.com/cenkalti/backoff/v5 v5.0.3 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/cilium/ebpf v0.19.0 // indirect github.com/containerd/errdefs v1.0.0 // indirect github.com/containerd/errdefs/pkg v0.3.0 // indirect github.com/containerd/log v0.1.0 // indirect @@ -237,6 +239,7 @@ require ( github.com/googleapis/gax-go/v2 v2.24.1 // indirect github.com/goreleaser/chglog v0.7.4 // indirect github.com/gorilla/handlers v1.5.2 // indirect + github.com/grafana/pyroscope-go/godeltaprof v0.1.11 // indirect github.com/hashicorp/errwrap v1.1.0 // indirect github.com/hashicorp/go-cleanhttp v0.5.2 // indirect github.com/hashicorp/go-retryablehttp v0.7.8 // indirect @@ -260,7 +263,7 @@ require ( github.com/josharian/intern v1.0.0 // indirect github.com/kelseyhightower/envconfig v1.4.0 // indirect github.com/kevinburke/ssh_config v1.4.0 // indirect - github.com/klauspost/compress v1.18.3 // indirect + github.com/klauspost/compress v1.18.7 // indirect github.com/klauspost/cpuid/v2 v2.3.0 // indirect github.com/koron/go-ssdp v0.0.4 // indirect github.com/kr/fs v0.1.0 // indirect @@ -366,7 +369,7 @@ replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-2 replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0 -replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260902163841-4a71f7b1d9e1 +replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260928122122-d1db5fa35318 tool ( github.com/goreleaser/chglog/cmd/chglog diff --git a/go.sum b/go.sum index c92b377bc..849ec8433 100644 --- a/go.sum +++ b/go.sum @@ -106,6 +106,8 @@ github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0= github.com/c-robinson/iplib v1.0.3 h1:NG0UF0GoEsrC1/vyfX1Lx2Ss7CySWl3KqqXh3q4DdPU= github.com/c-robinson/iplib v1.0.3/go.mod h1:i3LuuFL1hRT5gFpBRnEydzw8R6yhGkF4szNDIbF8pgo= +github.com/caarlos0/env/v11 v11.4.1 h1:fYwH0sWEsBSMPG7t4e/PEfTFzrWrpjyygXyUnWiSwEw= +github.com/caarlos0/env/v11 v11.4.1/go.mod h1:qupehSf/Y0TUTsxKywqRt/vJjN5nz6vauiYEUUr8P4U= github.com/caddyserver/certmagic v0.21.3 h1:pqRRry3yuB4CWBVq9+cUqu+Y6E2z8TswbhNx1AZeYm0= github.com/caddyserver/certmagic v0.21.3/go.mod h1:Zq6pklO9nVRl3DIFUw9gVUfXKdpc/0qwTUAQMBlfgtI= github.com/caddyserver/zerossl v0.1.3 h1:onS+pxp3M8HnHpN5MMbOMyNjmTheJyWRaZYwn+YTAyA= @@ -327,6 +329,10 @@ github.com/gorilla/handlers v1.5.2 h1:cLTUSsNkgcwhgRqvCNmdbRWG0A3N4F+M2nWKdScwyE github.com/gorilla/handlers v1.5.2/go.mod h1:dX+xVpaxdSw+q0Qek8SSsl3dfMk3jNddUkMzo0GtH0w= github.com/gorilla/mux v1.8.1 h1:TuBL49tXwgrFYWhqrNgrUNEY92u81SPhu7sTdzQEiWY= github.com/gorilla/mux v1.8.1/go.mod h1:AKf9I4AEqPTmMytcMc0KkNouC66V3BtZ4qD5fmWSiMQ= +github.com/grafana/pyroscope-go v1.4.2 h1:0LW5HrUJXgGr9zF5gITP/HaFXN9/LsMiwlgVJAK75l0= +github.com/grafana/pyroscope-go v1.4.2/go.mod h1:Ej13Jr05rRJrjWvrrFhfh6gGYXtfibuukOs3Tl3Y7QQ= +github.com/grafana/pyroscope-go/godeltaprof v0.1.11 h1:el5LYpXissAiCKZ5/6yjlr6mhYVV6Cp5lahTocxraXM= +github.com/grafana/pyroscope-go/godeltaprof v0.1.11/go.mod h1:jl1V8M4cWsXciROCPIDDG7CtjSjT/ECbp6eLVuMxYRI= github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357 h1:Fkzd8ktnpOR9h47SXHe2AYPwelXLH2GjGsjlAloiWfo= github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357/go.mod h1:w9Y7gY31krpLmrVU5ZPG9H7l9fZuRu5/3R3S3FMtVQ4= github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5ukBEgSGXEN89zeH1Jo= @@ -401,8 +407,6 @@ github.com/jonboulle/clockwork v0.5.0 h1:Hyh9A8u51kptdkR+cqRpT1EebBwTn1oK9YfGYbd github.com/jonboulle/clockwork v0.5.0/go.mod h1:3mZlmanh0g2NDKO5TWZVJAfofYk64M7XN3SzBPjZF60= github.com/josharian/intern v1.0.0 h1:vlS4z54oSdjm0bgjRigI+G1HpF+tI+9rE5LLzOg8HmY= github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y= -github.com/jsimonetti/rtnetlink/v2 v2.0.1 h1:xda7qaHDSVOsADNouv7ukSuicKZO7GgVUCXxpaIEIlM= -github.com/jsimonetti/rtnetlink/v2 v2.0.1/go.mod h1:7MoNYNbb3UaDHtF8udiJo/RH6VsTKP1pqKLUTVCvToE= github.com/json-iterator/go v1.1.7/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= github.com/jtolds/gls v4.20.0+incompatible h1:xdiiI2gbIgH/gLH7ADydsJ1uDOEzR8yvV7C0MuV77Wo= github.com/jtolds/gls v4.20.0+incompatible/go.mod h1:QJZ7F/aHp+rZTRtaJ1ow/lLfFfVYBRgL+9YlvaHOwJU= @@ -413,8 +417,8 @@ github.com/kevinburke/ssh_config v1.4.0 h1:6xxtP5bZ2E4NF5tuQulISpTO2z8XbtH8cg1PW github.com/kevinburke/ssh_config v1.4.0/go.mod h1:q2RIzfka+BXARoNexmF9gkxEX7DmvbW9P4hIVx2Kg4M= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= -github.com/klauspost/compress v1.18.3 h1:9PJRvfbmTabkOX8moIpXPbMMbYN60bWImDDU7L+/6zw= -github.com/klauspost/compress v1.18.3/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4= +github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw= +github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/cpuid/v2 v2.0.12/go.mod h1:g2LTdtYhdyuGPqyWyv7qRAmj1WBqxuObKfj5c0PQa7c= github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y= github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0= @@ -521,8 +525,8 @@ github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9ax github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ= -github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260902163841-4a71f7b1d9e1 h1:n5aXV/U6I9bLc+yWN088TyVR4OfF64Gy+L6Hrffc+n4= -github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260902163841-4a71f7b1d9e1/go.mod h1:/6QR46/nhGCSADHbS++XtDb9dkTnenTHlGskTPRo9S0= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260928122122-d1db5fa35318 h1:Qv2jYeucRkuqkkM9aAVFcO9avmSfPEoB+gxKKKuI4vA= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260928122122-d1db5fa35318/go.mod h1:62UsqQRqanuCbFi6xElIw3/2X+yGF+TKlyiWImA5taU= github.com/netbirdio/wireguard-go v0.0.0-20260914123147-8bf8fa968f1a h1:Nt8BgkTkI56LGBPPEBywM406MVKJDqeDIVdgsZyYs80= github.com/netbirdio/wireguard-go v0.0.0-20260914123147-8bf8fa968f1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw= github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A= diff --git a/idp/dex/provider.go b/idp/dex/provider.go index f40b96a58..95c511c01 100644 --- a/idp/dex/provider.go +++ b/idp/dex/provider.go @@ -737,6 +737,10 @@ func (p *Provider) UpdateUserPassword(ctx context.Context, userID string, oldPas return fmt.Errorf("failed to update password: %w", err) } + if err := p.storage.DeleteAuthSession(ctx, user.UserID, server.LocalConnector); err != nil && !errors.Is(err, storage.ErrNotFound) { + p.logger.Error("failed to revoke local session after password change", "error", err) + } + return nil } diff --git a/infrastructure_files/configure.sh b/infrastructure_files/configure.sh index 92252d0b3..ce1a041e6 100755 --- a/infrastructure_files/configure.sh +++ b/infrastructure_files/configure.sh @@ -1,14 +1,14 @@ #!/bin/bash set -e -if ! which curl >/dev/null 2>&1; then +if ! command -v curl >/dev/null 2>&1; then echo "This script uses curl fetch OpenID configuration from IDP." echo "Please install curl and re-run the script https://curl.se/" echo "" exit 1 fi -if ! which jq >/dev/null 2>&1; then +if ! command -v jq >/dev/null 2>&1; then echo "This script uses jq to load OpenID configuration from IDP." echo "Please install jq and re-run the script https://stedolan.github.io/jq/" echo "" @@ -18,13 +18,13 @@ fi source setup.env source base.setup.env -if ! which envsubst >/dev/null 2>&1; then +if ! command -v envsubst >/dev/null 2>&1; then echo "envsubst is needed to run this script" if [[ $(uname) == "Darwin" ]]; then echo "you can install it with homebrew (https://brew.sh):" echo "brew install gettext" else - if which apt-get >/dev/null 2>&1; then + if command -v apt-get >/dev/null 2>&1; then echo "you can install it by running" echo "apt-get update && apt-get install gettext-base" else diff --git a/infrastructure_files/getting-started-enterprise.sh b/infrastructure_files/getting-started-enterprise.sh index e88436e84..344bd2fed 100755 --- a/infrastructure_files/getting-started-enterprise.sh +++ b/infrastructure_files/getting-started-enterprise.sh @@ -6,11 +6,19 @@ set -o pipefail # NetBird Enterprise — Getting Started # Single-node bootstrap for a self-hosted NetBird Enterprise stack with the # embedded identity provider. Owner is created via first-login flow. +# Add features to an existing install with --enable-proxy or --enable-traffic-events. SED_STRIP_PADDING='s/=//g' NETBIRD_EULA_URL="https://netbird.io/self-hosted-EULA" +STACK_FILES=(.env docker-compose.yml config.yaml) + +# Host directory of a custom TLS certificate mounted at /certs, see +# https://docs.netbird.io/selfhosted/enterprise/getting-started#appendix-using-a-custom-tls-certificate +CUSTOM_TLS_CERTS="" +PROXY_TOKEN_ID="" + # Static IP for Traefik inside the compose bridge network. The management # server trusts X-Forwarded-* headers from this address only. TRAEFIK_IP="172.30.0.10" @@ -44,6 +52,28 @@ check_openssl() { fi } +die() { + echo "$1" > /dev/stderr + exit 1 +} + +# env_get KEY [DEFAULT] prints KEY's value from .env, or DEFAULT if unset. +env_get() { + local value + value=$(sed -n "s/^$1=//p" .env | tail -n 1) + echo "${value:-$2}" +} + +# merge_env upserts the KEY=VALUE lines from stdin into .env, in place. +merge_env() { + local merged + merged=$(awk -F= 'NR == FNR { v[$1] = $0; o[++n] = $1; next } + $1 in v { print v[$1]; delete v[$1]; next } + { print } + END { for (i = 1; i <= n; i++) if (o[i] in v) print v[o[i]] }' - .env) + printf '%s\n' "$merged" > .env +} + rand_secret() { openssl rand -base64 32 | sed "$SED_STRIP_PADDING" } @@ -171,6 +201,15 @@ read_yes_no() { esac } +read_crowdsec_option() { + echo "" + echo "CrowdSec:" + echo " Checks client IPs against a community threat intelligence database and" + echo " blocks known malicious sources before they reach services exposed through" + echo " the proxy. Adds a CrowdSec container to the stack." + NETBIRD_CROWDSEC=$(read_yes_no "Enable CrowdSec" "n") +} + # Gate the install on explicit acceptance of the NetBird On-Premise EULA. require_eula_acceptance() { cat > /dev/stderr < /dev/null && return 0 + sleep 2 + done + return 1 +} + +admin_token() { + $DOCKER_COMPOSE_COMMAND run --rm --no-deps -T netbird-server admin token "$@" --config /etc/netbird/config.yaml +} + +# revoke_proxy_token revokes the token this run minted if the proxy never started. +# On failure the ID is kept, so rollback retries it. +revoke_proxy_token() { + [[ -n "$PROXY_TOKEN_ID" ]] || return 0 + if admin_token revoke "$PROXY_TOKEN_ID" > /dev/null; then + PROXY_TOKEN_ID="" + return 0 + fi + echo "Could not revoke the unused proxy token ${PROXY_TOKEN_ID}. Revoke it with:" > /dev/stderr + echo " $DOCKER_COMPOSE_COMMAND run --rm netbird-server admin token revoke ${PROXY_TOKEN_ID} --config /etc/netbird/config.yaml" > /dev/stderr +} + +# start_proxy mints the proxy token and CrowdSec bouncer key, then starts the proxy. +start_proxy() { + local out token key + echo "Creating the proxy access token ..." + out=$(admin_token create --name default-proxy) || true + token=$(awk '/^Token:/ {print $2}' <<< "$out") + PROXY_TOKEN_ID=$(awk '/^Token ID:/ {print $3}' <<< "$out") + [[ -n "$token" ]] || die "Could not create the proxy access token. Check the netbird-server logs, then re-run with --enable-proxy." + + if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then + echo "Registering the CrowdSec bouncer ..." + if wait_crowdsec; then + # "add" fails if an earlier attempt already registered the bouncer. + $DOCKER_COMPOSE_COMMAND exec -T crowdsec cscli bouncers delete netbird-proxy &> /dev/null || true + key=$($DOCKER_COMPOSE_COMMAND exec -T crowdsec cscli bouncers add netbird-proxy -o raw) || true + fi + if [[ -z "$key" ]]; then + revoke_proxy_token + die "Could not register the CrowdSec bouncer. Check the crowdsec logs, then re-run with --enable-proxy." + fi + fi + + # A stored token marks the proxy as set up, so it is cleared again on failure. + { + echo "NETBIRD_PROXY_TOKEN=${token}" + if [[ -n "$key" ]]; then echo "NETBIRD_CROWDSEC_BOUNCER_KEY=${key}"; fi + } | merge_env + if ! $DOCKER_COMPOSE_COMMAND up -d proxy; then + revoke_proxy_token + echo "NETBIRD_PROXY_TOKEN=" | merge_env + die "Could not start the proxy. Check the proxy logs, then re-run with --enable-proxy." + fi + PROXY_TOKEN_ID="" +} + +print_proxy_notes() { + echo "" + echo "NetBird Proxy:" + echo " Every domain other than ${NETBIRD_DOMAIN} is passed through to the proxy," + echo " which issues its own TLS certificates. Point proxy domains at this host:" + echo "" + echo " *.${NETBIRD_DOMAIN} CNAME ${NETBIRD_DOMAIN}" + echo "" + echo " Open 51820/udp (optional) for peer-to-peer proxy connections." + if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then + echo " CrowdSec is running. Enable it per service in the dashboard under Access Control." + fi +} + init_environment() { check_openssl DOCKER_COMPOSE_COMMAND=$(check_docker_compose) if [[ -f .env ]] || [[ -f docker-compose.yml ]] || [[ -f config.yaml ]]; then echo "Generated files already exist in $(pwd)." + echo "To add the proxy or traffic events to this installation, re-run with" + echo "--enable-proxy or --enable-traffic-events." + echo "" echo "If you want to reinitialize the environment, please remove them first:" echo " $DOCKER_COMPOSE_COMMAND down --volumes # removes all containers and volumes" - echo " rm -f .env docker-compose.yml config.yaml" + echo " rm -rf .env docker-compose.yml config.yaml traefik" echo "Be aware this will remove all data from the database." exit 1 fi @@ -341,6 +465,16 @@ init_environment() { echo " See https://docs.netbird.io/manage/activity/traffic-events-logging" NETBIRD_TRAFFIC_FLOW=$(read_yes_no "Enable traffic flow" "n") + echo "" + echo "NetBird Proxy:" + echo " Exposes selected resources from your NetBird network to the internet." + echo " You choose which resources are exposed from the dashboard." + NETBIRD_PROXY=$(read_yes_no "Enable the NetBird Proxy" "n") + NETBIRD_CROWDSEC="no" + if [[ "$NETBIRD_PROXY" == "yes" ]]; then + read_crowdsec_option + fi + echo "" NETBIRD_DOMAIN=$(read_nb_domain) @@ -364,6 +498,8 @@ init_environment() { echo "" echo "Selected:" echo " Traffic flow: ${NETBIRD_TRAFFIC_FLOW}" + echo " Proxy: ${NETBIRD_PROXY}" + echo " CrowdSec: ${NETBIRD_CROWDSEC}" echo " Domain: ${NETBIRD_DOMAIN}" echo " ACME email: ${NETBIRD_LETSENCRYPT_EMAIL}" echo "" @@ -371,9 +507,9 @@ init_environment() { install -m 600 /dev/null .env render_env >> .env render_docker_compose > docker-compose.yml - - if [[ -z "${NETBIRD_LICENSE_SERVER_BASE_URL:-}" ]]; then - sed -i.bak '/NETBIRD_LICENSE_SERVER_BASE_URL/d' docker-compose.yml && rm -f docker-compose.yml.bak + mkdir -p traefik + if [[ "$NETBIRD_PROXY" == "yes" ]]; then + render_traefik_proxy > traefik/proxy.yaml fi install -m 600 /dev/null config.yaml render_config_yaml >> config.yaml @@ -390,11 +526,16 @@ init_environment() { echo "" echo "Starting remaining services ..." - $DOCKER_COMPOSE_COMMAND up -d + up_all_but_proxy echo "" wait_for_license_verdict + if [[ "$NETBIRD_PROXY" == "yes" ]]; then + echo "" + start_proxy + fi + echo "" echo "Done." echo "" @@ -402,6 +543,9 @@ init_environment() { echo "" echo "Open the dashboard in a browser to complete the first-login owner setup." echo "All configuration and secrets are stored (mode 600) in $(pwd)/.env" + if [[ "$NETBIRD_PROXY" == "yes" ]]; then + print_proxy_notes + fi echo "" echo "Tail logs:" echo " cd $(pwd) && $DOCKER_COMPOSE_COMMAND logs -f netbird-server traefik" @@ -413,6 +557,148 @@ init_environment() { fi } +# service_block NAME prints a service's definition from the compose file on stdin. +service_block() { + local name="$1" + awk -v s=" ${name}:" '$0 == s { p = 1; print; next } p && (/^[^ ]/ || /^ [^ ]/) { exit } p' +} + +# enable_features adds the proxy and/or traffic events to the install in the +# current directory, restoring the backed-up files if any step fails. +enable_features() { + local want_proxy="$1" want_flow="$2" f compose + DOCKER_COMPOSE_COMMAND=$(check_docker_compose) + + for f in "${STACK_FILES[@]}"; do + [[ -f "$f" ]] || die "$f not found in $(pwd). Run this from an existing installation directory." + done + grep -q '^# Generated by getting-started-enterprise.sh' .env || die ".env was not generated by getting-started-enterprise.sh." + # Installs from before the move to Traefik run Caddy and can't be re-rendered. + [[ -n "$(env_get NETBIRD_TRAEFIK_IP)" ]] || die "This installation predates the Traefik layout and can't be updated in place." + + NETBIRD_DOMAIN=$(env_get NETBIRD_DOMAIN) + NETBIRD_LICENSE_SERVER_BASE_URL=$(env_get NETBIRD_LICENSE_SERVER_BASE_URL) + NETBIRD_TRAFFIC_FLOW=$(env_get NETBIRD_TRAFFIC_FLOW_ENABLED no) + NETBIRD_PROXY=$(env_get NETBIRD_PROXY_ENABLED no) + NETBIRD_CROWDSEC=$(env_get NETBIRD_CROWDSEC_ENABLED no) + # A custom certificate counts only once Traefik mounts it and ACME is already gone, + # so the re-render never removes a working Let's Encrypt setup. + local traefik_block + traefik_block=$(service_block traefik < docker-compose.yml) + if ! grep -q certificatesresolvers <<< "$traefik_block"; then + CUSTOM_TLS_CERTS=$(awk '/:\/certs:ro$/ { sub(/^ *- /, ""); sub(/:\/certs:ro$/, ""); print; exit }' <<< "$traefik_block") + fi + + if [[ "$want_flow" == "yes" && "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then + echo "Traffic events are already enabled." + want_flow="no" + fi + # No token means an earlier proxy setup failed, so let it run again. + if [[ "$want_proxy" == "yes" && -n "$(env_get NETBIRD_PROXY_TOKEN)" ]]; then + echo "The NetBird Proxy is already enabled." + want_proxy="no" + fi + if [[ "$want_flow" == "no" && "$want_proxy" == "no" ]]; then + exit 0 + fi + + if [[ "$want_flow" == "yes" ]]; then + NETBIRD_TRAFFIC_FLOW="yes" + fi + # A retry keeps the CrowdSec choice made at install time. + if [[ "$want_proxy" == "yes" && "$NETBIRD_PROXY" != "yes" ]]; then + read_crowdsec_option + fi + if [[ "$want_proxy" == "yes" ]]; then + NETBIRD_PROXY="yes" + fi + + compose=$(render_docker_compose) + echo "" + echo "Changes to docker-compose.yml:" + printf '%s\n' "$compose" | diff -u docker-compose.yml - || true + + # Lines the new file drops are most likely local edits, so don't default to applying. + local lost s restarts="" apply="y" + lost=$(printf '%s\n' "$compose" | awk 'NR == FNR { keep[$0]; next } !($0 in keep)' - docker-compose.yml) + if [[ -n "$lost" ]]; then + echo "" + echo "These lines are not in the new docker-compose.yml and will be lost:" + printf '%s\n' "$lost" + apply="n" + fi + for s in $($DOCKER_COMPOSE_COMMAND config --services); do + if [[ "$(service_block "$s" < docker-compose.yml)" != "$(printf '%s\n' "$compose" | service_block "$s")" ]] \ + || [[ "$s" == "netbird-server" && "$want_flow" == "yes" ]]; then + restarts+=" $s" + fi + done + echo "" + if [[ -n "$restarts" ]]; then + echo "These services will restart:${restarts}" + fi + if [[ "$(read_yes_no "Apply these changes?" "$apply")" != "yes" ]]; then + echo "Aborted." + exit 0 + fi + + BACKUP_SUFFIX=".bak.$(date -u +%Y%m%d%H%M%S)" + for f in "${STACK_FILES[@]}" traefik/proxy.yaml; do + if [[ -f "$f" ]]; then cp -p "$f" "$f$BACKUP_SUFFIX"; fi + done + trap rollback EXIT + + printf '%s\n' "$compose" > docker-compose.yml + { + echo "NETBIRD_TRAFFIC_FLOW_ENABLED=${NETBIRD_TRAFFIC_FLOW}" + echo "NETBIRD_PROXY_ENABLED=${NETBIRD_PROXY}" + echo "NETBIRD_CROWDSEC_ENABLED=${NETBIRD_CROWDSEC}" + if [[ "$want_flow" == "yes" ]]; then render_env_flow; fi + if [[ "$want_proxy" == "yes" ]]; then render_env_proxy; fi + } | merge_env + if [[ "$want_flow" == "yes" ]] && ! grep -q '^ trafficFlow:' config.yaml; then + render_config_flow >> config.yaml + fi + mkdir -p traefik + if [[ "$want_proxy" == "yes" ]]; then + render_traefik_proxy > traefik/proxy.yaml + fi + + up_all_but_proxy + if [[ "$want_flow" == "yes" ]]; then + # Compose does not notice changes to the bind-mounted config.yaml. + $DOCKER_COMPOSE_COMMAND restart netbird-server + fi + if [[ "$want_proxy" == "yes" ]]; then + start_proxy + fi + trap - EXIT + + echo "" + echo "Done. The previous files are kept with the ${BACKUP_SUFFIX} suffix." + if [[ "$want_flow" == "yes" ]]; then + echo "" + echo "Traffic events still have to be turned on from the dashboard settings." + echo " See https://docs.netbird.io/manage/activity/traffic-events-logging" + fi + if [[ "$want_proxy" == "yes" ]]; then + print_proxy_notes + fi +} + +rollback() { + local f + echo "" > /dev/stderr + echo "Enabling failed. Restoring the previous configuration ..." > /dev/stderr + revoke_proxy_token + # Files without a backup were created by this run. + for f in "${STACK_FILES[@]}" traefik/proxy.yaml; do + if [[ -f "$f$BACKUP_SUFFIX" ]]; then cp -p "$f$BACKUP_SUFFIX" "$f"; else rm -f "$f"; fi + done + $DOCKER_COMPOSE_COMMAND up -d --remove-orphans + $DOCKER_COMPOSE_COMMAND restart netbird-server +} + # ------------------------------------------------------------------ # Renderers # ------------------------------------------------------------------ @@ -427,8 +713,10 @@ NETBIRD_EULA_ACCEPTED=yes NETBIRD_EULA_ACCEPTED_AT=${NETBIRD_EULA_ACCEPTED_AT} NETBIRD_EULA_URL=${NETBIRD_EULA_URL} -# Features (set by the script; don't edit without re-running) +# Features (change with --enable-proxy or --enable-traffic-events, not by hand) NETBIRD_TRAFFIC_FLOW_ENABLED=${NETBIRD_TRAFFIC_FLOW} +NETBIRD_PROXY_ENABLED=${NETBIRD_PROXY} +NETBIRD_CROWDSEC_ENABLED=${NETBIRD_CROWDSEC} # Domain NETBIRD_DOMAIN=${NETBIRD_DOMAIN} @@ -444,10 +732,11 @@ NETBIRD_SERVER_TAG=${NETBIRD_SERVER_TAG:-latest} EOF if [[ "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then - cat < /dev/stderr; exit 1 ;; + esac + shift + done + + if [[ "$enable_proxy" == "no" && "$enable_flow" == "no" ]]; then + init_environment + else + enable_features "$enable_proxy" "$enable_flow" fi } -init_environment +main "$@" diff --git a/infrastructure_files/migrate-to-enterprise.sh b/infrastructure_files/migrate-to-enterprise.sh index 2b10250c9..8abde2f68 100755 --- a/infrastructure_files/migrate-to-enterprise.sh +++ b/infrastructure_files/migrate-to-enterprise.sh @@ -220,6 +220,21 @@ detect_exposed_address() { yq eval '.server.exposedAddress // ""' "$CONFIG_YAML_HOST" } +detect_relay_auth_secret() { + local secret="" + local external_relay_count + external_relay_count=$(yq eval '(.server.relays.addresses // []) | length' "$CONFIG_YAML_HOST") + + if (( external_relay_count > 0 )); then + secret=$(yq eval '.server.relays.secret // ""' "$CONFIG_YAML_HOST") + fi + if [[ -z "$secret" ]] || [[ "$secret" == "null" ]]; then + secret=$(yq eval '.server.authSecret // ""' "$CONFIG_YAML_HOST") + fi + + printf '%s' "$secret" +} + # The engine is a config.yaml-only setting — there is no env override for it # (combined/cmd/root.go reads it from YAML and derives the env vars), so # config.yaml is authoritative. Absent means the sqlite default. @@ -945,11 +960,10 @@ init_migration() { if [[ "$MIGRATE_POSTGRES" == "yes" ]] || [[ "$EXISTING_POSTGRES" == "yes" ]]; then ENABLE_FLOW=$(read_yes_no "Step 3: Enable traffic flow? (requires Postgres)" "n") if [[ "$ENABLE_FLOW" == "yes" ]]; then - # Auth secret MUST match server.authSecret from config.yaml - NB_FLOW_AUTH_SECRET=$(yq eval '.server.authSecret // ""' "$CONFIG_YAML_HOST") + NB_FLOW_AUTH_SECRET=$(detect_relay_auth_secret) if [[ -z "$NB_FLOW_AUTH_SECRET" ]] || [[ "$NB_FLOW_AUTH_SECRET" == "null" ]]; then - echo "Could not read server.authSecret from $CONFIG_YAML_HOST." > /dev/stderr - echo "Flow receiver auth must match the combined server's authSecret." > /dev/stderr + echo "Could not resolve the Relay auth secret from $CONFIG_YAML_HOST." > /dev/stderr + echo "Set server.relays.secret for external Relays or server.authSecret for the local Relay." > /dev/stderr exit 1 fi diff --git a/integration_tests/management/network_map_db/sqlite_test_store.go b/integration_tests/management/network_map_db/sqlite_test_store.go index 1c70c93d4..622ac1c83 100644 --- a/integration_tests/management/network_map_db/sqlite_test_store.go +++ b/integration_tests/management/network_map_db/sqlite_test_store.go @@ -9,8 +9,8 @@ import ( "strings" networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite" + nbdb "github.com/netbirdio/netbird/management/internals/shared/db" gormstore "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" log "github.com/sirupsen/logrus" "gorm.io/driver/sqlite" "gorm.io/gorm" @@ -28,7 +28,11 @@ func createSqliteTestStore(baseData string) (*networkmap_sqlite.SqliteStore, fun if err != nil { log.Fatalf("error initializing db: %s", err.Error()) } - _, err = gormstore.NewSqlStore(context.TODO(), db, types.SqliteStoreEngine, nil, false) + conn, err := nbdb.NewConn(context.TODO(), db, nbdb.SqliteStoreEngine, nil) + if err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + } + _, err = gormstore.NewSqlStore(context.TODO(), conn, nil, false) if err != nil { log.Fatalf("error initializing db: %s", err.Error()) } diff --git a/management/cmd/proxy/proxy.go b/management/cmd/proxy/proxy.go index 73f83b3d6..1186c8d61 100644 --- a/management/cmd/proxy/proxy.go +++ b/management/cmd/proxy/proxy.go @@ -10,6 +10,7 @@ import ( "io" "strings" "text/tabwriter" + "unicode" "github.com/spf13/cobra" @@ -68,8 +69,8 @@ func runDisconnectAll(ctx context.Context, s store.Store, out io.Writer, in io.R toDisconnect := 0 w := tabwriter.NewWriter(out, 0, 0, 2, ' ', 0) - _, _ = fmt.Fprintln(w, "ID\tCLUSTER\tIP\tACCOUNT\tSTATUS\tLAST SEEN") - _, _ = fmt.Fprintln(w, "--\t-------\t--\t-------\t------\t---------") + _, _ = fmt.Fprintln(w, "ID\tCLUSTER\tIP\tVERSION\tACCOUNT\tSTATUS\tLAST SEEN") + _, _ = fmt.Fprintln(w, "--\t-------\t--\t-------\t-------\t------\t---------") for _, p := range proxies { if p.Status != rpproxy.StatusDisconnected { @@ -80,11 +81,16 @@ func runDisconnectAll(ctx context.Context, s store.Store, out io.Writer, in io.R if p.AccountID != nil { account = *p.AccountID } + version := "-" + if p.Version != "" { + version = sanitizeReportedValue(p.Version) + } - _, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\n", - p.ID, + _, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\t%s\n", + sanitizeReportedValue(p.ID), p.ClusterAddress, p.IPAddress, + version, account, p.Status, p.LastSeen.Format("2006-01-02 15:04:05"), @@ -139,3 +145,16 @@ func confirmDisconnectAll(out io.Writer, in io.Reader) (bool, error) { return strings.EqualFold(strings.TrimSpace(scanner.Text()), disconnectAllConfirmation), nil } + +// sanitizeReportedValue replaces non-printable characters in a value the proxy +// reports about itself. Both the id and the version arrive unvalidated over +// gRPC, so a tab would forge a column, a carriage return or ANSI escape would +// redraw the operator's terminal, and U+202E would reverse the rest of the line. +func sanitizeReportedValue(s string) string { + return strings.Map(func(r rune) rune { + if unicode.IsPrint(r) { + return r + } + return '\uFFFD' + }, s) +} diff --git a/management/cmd/proxy/proxy_test.go b/management/cmd/proxy/proxy_test.go index ff0dc8119..6e3cd0c01 100644 --- a/management/cmd/proxy/proxy_test.go +++ b/management/cmd/proxy/proxy_test.go @@ -35,6 +35,7 @@ func seedProxies(t *testing.T, ctx context.Context, s store.Store) { SessionID: "session-1", ClusterAddress: "cluster-a.example.com", IPAddress: "10.0.0.1", + Version: "0.60.0", LastSeen: time.Now(), Status: rpproxy.StatusConnected, }, @@ -89,6 +90,7 @@ func TestRunDisconnectAllWithConfirmation(t *testing.T) { require.Contains(t, output, "proxy-2") require.Contains(t, output, "proxy-3") require.Contains(t, output, "cluster-a.example.com") + require.Contains(t, output, "0.60.0") require.Contains(t, output, "account-1") require.Contains(t, output, "Type \"disconnect all proxies\" to continue") require.Contains(t, output, "Force-marked 2 of 3 reverse proxy instance(s) as disconnected.") @@ -178,3 +180,40 @@ func TestRunDisconnectAllEmpty(t *testing.T) { require.NoError(t, runDisconnectAll(ctx, s, &out, strings.NewReader(""), false, false)) require.Contains(t, out.String(), "No reverse proxy instances found.") } + +func TestRunDisconnectAllEscapesProxyReportedFields(t *testing.T) { + ctx := context.Background() + s := newTestStore(t) + + // A proxy reports its own id and version on connect, so both reach this + // listing unvalidated. Carriage returns, tabs and ANSI escapes would let + // a malicious proxy redraw the table or forge a row on the operator's + // terminal; U+202E would reverse the rendering of the rest of the line. + require.NoError(t, s.SaveProxy(ctx, &rpproxy.Proxy{ + ID: "proxy-\r\x1b[2Kevil", + SessionID: "session-1", + ClusterAddress: "cluster-a.example.com", + IPAddress: "10.0.0.1", + Version: "0.60.0\tfake\rcolumn\u202e", + LastSeen: time.Now(), + Status: rpproxy.StatusConnected, + })) + + var out bytes.Buffer + require.NoError(t, runDisconnectAll(ctx, s, &out, strings.NewReader(disconnectAllConfirmation+"\n"), true, false)) + + output := out.String() + for _, forbidden := range []string{"\r", "\x1b", "\u202e"} { + require.NotContains(t, output, forbidden, "listing must not carry proxy-reported control characters") + } + // The table has one data row; a smuggled tab would add a phantom column. + var dataRow string + for _, line := range strings.Split(output, "\n") { + if strings.Contains(line, "evil") { + dataRow = line + } + } + require.NotEmpty(t, dataRow, "listing should still show the proxy row") + require.NotContains(t, dataRow, "\t", "tabwriter output should not carry a smuggled column separator") + require.Contains(t, dataRow, "0.60.0", "the printable part of the version should survive") +} diff --git a/management/internals/controllers/network_map/controller/repository.go b/management/internals/controllers/network_map/controller/repository.go index 5c3195f16..c11af0b69 100644 --- a/management/internals/controllers/network_map/controller/repository.go +++ b/management/internals/controllers/network_map/controller/repository.go @@ -44,7 +44,7 @@ func (r *repository) GetAccountNetwork(ctx context.Context, accountID string) (* } func (r *repository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) { - return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") } func (r *repository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) { diff --git a/management/internals/controllers/network_map/nmaptest/runner.go b/management/internals/controllers/network_map/nmaptest/runner.go index 0d6ac9c18..ffce6483e 100644 --- a/management/internals/controllers/network_map/nmaptest/runner.go +++ b/management/internals/controllers/network_map/nmaptest/runner.go @@ -245,7 +245,7 @@ func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkma peerGroups := maps.Keys(nmData.GetPeerGroups(peerID)) resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil, dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort) - res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain) + res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain, false) require.NoError(t, err, "expand envelope") return res.NetworkMap default: diff --git a/management/internals/modules/agentnetwork/accesslog_cleanup_realstore_test.go b/management/internals/modules/agentnetwork/accesslog_cleanup_realstore_test.go new file mode 100644 index 000000000..0e294a560 --- /dev/null +++ b/management/internals/modules/agentnetwork/accesslog_cleanup_realstore_test.go @@ -0,0 +1,76 @@ +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/store" + nbtypes "github.com/netbirdio/netbird/management/server/types" +) + +// TestCleanupAccessLogs_RealStore_DeletedAccount covers a deleted account's access logs. +// The sweep is driven by settings rows, which go with the account, so without a fallback +// those logs would never expire. They get the default retention instead. A live account +// can delete its own settings row, so "no settings" must not be mistaken for "deleted": +// that account's logs are left alone, as are those of an account that keeps logs forever. +func TestCleanupAccessLogs_RealStore_DeletedAccount(t *testing.T) { + ctx := context.Background() + s, cleanup, err := store.NewTestStoreFromSQL(ctx, "", t.TempDir()) + require.NoError(t, err, "real sqlite test store must come up") + defer cleanup() + + const ( + deletedAccountID = "acc-deleted" + keepAccountID = "acc-keep-forever" + noSettingsAccountID = "acc-live-no-settings" + ) + old := time.Now().UTC().AddDate(0, 0, -(types.DefaultAccessLogRetentionDays + 10)) + recent := time.Now().UTC().AddDate(0, 0, -1) + + require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: keepAccountID})) + require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: noSettingsAccountID})) + + keepSettings := types.DefaultSettings(keepAccountID) + keepSettings.Domain = "keep.gw.example.com" + keepSettings.AccessLogRetentionDays = 0 + require.NoError(t, s.SaveAgentNetworkSettings(ctx, keepSettings)) + + mkLog := func(id, accountID string, ts time.Time) { + t.Helper() + entry := &types.AgentNetworkAccessLog{ + ID: id, AccountID: accountID, ServiceID: "svc", Timestamp: ts, StatusCode: 200, Model: "gpt-4o", + } + groups := []types.AgentNetworkAccessLogGroup{{LogID: id, GroupID: "grp-eng", AccountID: accountID}} + require.NoError(t, s.CreateAgentNetworkAccessLog(ctx, entry, groups)) + } + mkLog("deleted-old", deletedAccountID, old) + mkLog("deleted-recent", deletedAccountID, recent) + mkLog("keep-old", keepAccountID, old) + mkLog("no-settings-old", noSettingsAccountID, old) + + m := &managerImpl{store: s} + m.cleanupAccessLogsOnce(ctx) + + logIDs := func(accountID string) []string { + t.Helper() + logs, _, err := s.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID, + types.AgentNetworkAccessLogFilter{Page: 1, PageSize: 50}) + require.NoError(t, err) + ids := make([]string, 0, len(logs)) + for _, l := range logs { + ids = append(ids, l.ID) + } + return ids + } + assert.Equal(t, []string{"deleted-recent"}, logIDs(deletedAccountID), + "a deleted account should have logs past the default retention swept") + assert.Equal(t, []string{"keep-old"}, logIDs(keepAccountID), + "an account with retention disabled should keep its old logs") + assert.Equal(t, []string{"no-settings-old"}, logIDs(noSettingsAccountID), + "a live account without a settings row should keep its old logs") +} diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index d1a5ebd7b..71ff53214 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -80,6 +80,9 @@ type Manager interface { ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error) StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int) + // RemoveAccountGateway drops the account's gateway mappings from the + // proxies. It runs as an account deletion hook. + RemoveAccountGateway(ctx context.Context, accountID string) error RecordConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds, tokensIn, tokensOut int64, costUSD float64) error RecordAccountBudgetUsage(ctx context.Context, accountID, userID string, groupIDs []string, tokensIn, tokensOut int64, costUSD float64) error RecordUsage(ctx context.Context, in RecordUsageInput) error @@ -1350,8 +1353,8 @@ func (m *managerImpl) scopeFilterToCaller(ctx context.Context, accountID, userID // StartAccessLogCleanup launches a background sweep that periodically deletes // each account's agent-network access-log rows older than that account's -// AccessLogRetentionDays. Usage records are never swept. A non-positive -// interval defaults to 24h. +// AccessLogRetentionDays, and the consumption counters of deleted accounts. +// Usage records are never swept. A non-positive interval defaults to 24h. func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int) { if cleanupIntervalHours <= 0 { cleanupIntervalHours = 24 @@ -1362,21 +1365,40 @@ func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupInterval ticker := time.NewTicker(interval) defer ticker.Stop() - m.cleanupAccessLogsOnce(ctx) // run once on startup + m.cleanupOnce(ctx) // run once on startup for { select { case <-ctx.Done(): return case <-ticker.C: - m.cleanupAccessLogsOnce(ctx) + m.cleanupOnce(ctx) } } }() } +func (m *managerImpl) cleanupOnce(ctx context.Context) { + m.cleanupAccessLogsOnce(ctx) + m.cleanupDeletedAccountConsumption(ctx) +} + +// cleanupDeletedAccountConsumption deletes the consumption counters of accounts +// that no longer exist. Best-effort: a failure is logged and retried next sweep. +func (m *managerImpl) cleanupDeletedAccountConsumption(ctx context.Context) { + deleted, err := m.store.DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx) + if err != nil { + log.WithContext(ctx).Warnf("agent-network consumption cleanup: %v", err) + return + } + if deleted > 0 { + log.WithContext(ctx).Infof("agent-network consumption cleanup: deleted %d counters of deleted accounts", deleted) + } +} + // cleanupAccessLogsOnce sweeps every account's expired access-log rows against -// its configured retention. Best-effort: a per-account failure is logged and -// the sweep continues. +// its configured retention. Deleted accounts, whose settings rows went with +// them, get the default retention. Best-effort: a per-account failure is +// logged and the sweep continues. func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) { settings, err := m.store.GetAllAgentNetworkSettings(ctx, store.LockingStrengthNone) if err != nil { @@ -1384,18 +1406,31 @@ func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) { return } for _, s := range settings { - if s.AccessLogRetentionDays <= 0 { - continue // keep indefinitely - } - cutoff := time.Now().UTC().AddDate(0, 0, -s.AccessLogRetentionDays) - deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, s.AccountID, cutoff) - if err != nil { - log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", s.AccountID, err) - continue - } - if deleted > 0 { - log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, s.AccountID, s.AccessLogRetentionDays) - } + m.cleanupAccountAccessLogs(ctx, s.AccountID, s.AccessLogRetentionDays) + } + + deleted, err := m.store.GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx) + if err != nil { + log.WithContext(ctx).Errorf("agent-network access-log cleanup: list deleted accounts: %v", err) + return + } + for _, accountID := range deleted { + m.cleanupAccountAccessLogs(ctx, accountID, types.DefaultAccessLogRetentionDays) + } +} + +func (m *managerImpl) cleanupAccountAccessLogs(ctx context.Context, accountID string, retentionDays int) { + if retentionDays <= 0 { + return // keep indefinitely + } + cutoff := time.Now().UTC().AddDate(0, 0, -retentionDays) + deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, accountID, cutoff) + if err != nil { + log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", accountID, err) + return + } + if deleted > 0 { + log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, accountID, retentionDays) } } @@ -1545,6 +1580,8 @@ func (*mockManager) GetUsageOverview(_ context.Context, _, _ string, _ types.Age func (*mockManager) StartAccessLogCleanup(_ context.Context, _ int) {} +func (*mockManager) RemoveAccountGateway(_ context.Context, _ string) error { return nil } + func (*mockManager) RecordConsumption(_ context.Context, _ string, _ types.ConsumptionDimension, _ string, _, _, _ int64, _ float64) error { return nil } diff --git a/management/internals/modules/agentnetwork/reconcile.go b/management/internals/modules/agentnetwork/reconcile.go index 69e684014..20d0bb42a 100644 --- a/management/internals/modules/agentnetwork/reconcile.go +++ b/management/internals/modules/agentnetwork/reconcile.go @@ -2,8 +2,10 @@ package agentnetwork import ( "context" + "fmt" log "github.com/sirupsen/logrus" + goproto "google.golang.org/protobuf/proto" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/server/types" @@ -81,18 +83,66 @@ func (m *managerImpl) reconcile(ctx context.Context, accountID string) { } m.reconcileMu.Unlock() - for _, entry := range creates { - entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED - m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster) + m.sendMappings(ctx, accountID, creates, proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED) + m.sendMappings(ctx, accountID, updates, proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED) + m.sendMappings(ctx, accountID, deletes, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED) +} + +// sendMappings sends each entry as updateType. It sends a copy: the entries' +// mappings are shared with reconcileCache, which another reconcile or +// RemoveAccountGateway may be reading, so they are never written. +func (m *managerImpl) sendMappings(ctx context.Context, accountID string, entries []syntheticMapping, updateType proto.ProxyMappingUpdateType) { + for _, entry := range entries { + update := goproto.Clone(entry.mapping).(*proto.ProxyMapping) + update.Type = updateType + m.proxyController.SendServiceUpdateToCluster(ctx, accountID, update, entry.cluster) } - for _, entry := range updates { - entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED - m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster) +} + +// RemoveAccountGateway tells the proxies to drop every mapping of the account's +// gateway, so a deleted account's proxy config, provider API keys included, does +// not linger in proxy memory until the next resync. It is an account deletion +// hook: it runs before the account's data is removed, the last point at which +// the mappings can be synthesised from the store. The cache alone would miss +// them, since it is per instance and empty after a restart. If the deletion +// then fails, the gateway stays down until the account's next change reconciles +// it back. +func (m *managerImpl) RemoveAccountGateway(ctx context.Context, accountID string) error { + if m.proxyController == nil { + return nil } - for _, entry := range deletes { - entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED - m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster) + + services, err := SynthesizeServices(ctx, m.store, accountID) + if err != nil { + return fmt.Errorf("synthesise agent network services: %w", err) } + oidcCfg := m.proxyController.GetOIDCValidationConfig() + removed := make(map[string]syntheticMapping, len(services)) + for _, svc := range services { + if svc == nil || svc.ID == "" { + continue + } + removed[svc.ID] = syntheticMapping{ + mapping: svc.ToProtoMapping(rpservice.Delete, "", oidcCfg), + cluster: svc.ProxyCluster, + } + } + + m.reconcileMu.Lock() + for id, entry := range m.reconcileCache[accountID] { + if _, ok := removed[id]; !ok { + removed[id] = entry + } + } + delete(m.reconcileCache, accountID) + m.reconcileMu.Unlock() + + entries := make([]syntheticMapping, 0, len(removed)) + for _, entry := range removed { + entries = append(entries, entry) + } + m.sendMappings(ctx, accountID, entries, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED) + return nil } // diffMappings classifies the previous→current transition for a single diff --git a/management/internals/modules/agentnetwork/reconcile_test.go b/management/internals/modules/agentnetwork/reconcile_test.go index ab3b08481..2cfea9828 100644 --- a/management/internals/modules/agentnetwork/reconcile_test.go +++ b/management/internals/modules/agentnetwork/reconcile_test.go @@ -2,6 +2,8 @@ package agentnetwork import ( "context" + "sync" + "sync/atomic" "testing" "go.uber.org/mock/gomock" @@ -12,6 +14,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/shared/management/proto" + "github.com/netbirdio/netbird/shared/management/status" ) func newReconcileMgr(t *testing.T, ctrl *gomock.Controller) (*managerImpl, *store.MockStore, *proxy.MockController) { @@ -287,3 +290,154 @@ func TestDiffMappings_RemovedServiceIsDeletedOnItsOwnCluster(t *testing.T) { assert.Equal(t, "brave-otter.gateway.example.com", deletes[0].cluster) } } + +// TestRemoveAccountGateway_EmitsRemovedFromStore — account deletion runs on an +// instance that may never have reconciled the account, so its cache is empty. +// The mappings are synthesised from the store, still intact before the delete, +// and each is sent as REMOVED to the cluster that serves it. +func TestRemoveAccountGateway_EmitsRemovedFromStore(t *testing.T) { + ctx := context.Background() + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl) + provider := newReconcileTestProvider() + policy := newReconcileTestPolicy(provider.ID, "grp-eng") + + expectReconcileSynthInputs(mockStore, ctx, []*types.Provider{provider}, []*types.Policy{policy}, []*types.Guardrail{}) + mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{}) + + var sent []*proto.ProxyMapping + mockProxy.EXPECT(). + SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "eu.proxy.netbird.io"). + Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) { + sent = append(sent, m) + }) + + require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1")) + + require.Len(t, sent, 1, "the account's one gateway mapping must be removed") + assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type, "the update must be a removal") + assert.Equal(t, "agent-net-svc-acct-1", sent[0].Id, "the removal must name the account's gateway service") +} + +// TestRemoveAccountGateway_AlsoRemovesCachedMappings — a mapping this instance +// last sent but the store no longer synthesises (here, one on another cluster) +// is removed too, and the account's cache entry is cleared. +func TestRemoveAccountGateway_AlsoRemovesCachedMappings(t *testing.T) { + ctx := context.Background() + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl) + mgr.reconcileCache["acct-1"] = map[string]syntheticMapping{ + "stale-svc": {mapping: &proto.ProxyMapping{Id: "stale-svc"}, cluster: "us.proxy.netbird.io"}, + } + + // Settings but no providers: the store synthesises nothing. + mockStore.EXPECT(). + GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1"). + Return(newReconcileTestSettings(), nil) + mockStore.EXPECT(). + GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1"). + Return([]*types.Provider{}, nil) + mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{}) + + var sent []*proto.ProxyMapping + mockProxy.EXPECT(). + SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "us.proxy.netbird.io"). + Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) { + sent = append(sent, m) + }) + + require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1")) + + require.Len(t, sent, 1, "the cached mapping must be removed from its own cluster") + assert.Equal(t, "stale-svc", sent[0].Id) + assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type) + mgr.reconcileMu.Lock() + _, present := mgr.reconcileCache["acct-1"] + mgr.reconcileMu.Unlock() + assert.False(t, present, "the deleted account's cache entry must be cleared") +} + +// TestRemoveAccountGateway_SynthFailureAbortsDeletion — if the mappings cannot +// be read, nothing is sent and the error is returned, which as an account +// deletion hook keeps the account rather than leaving its gateway running. +func TestRemoveAccountGateway_SynthFailureAbortsDeletion(t *testing.T) { + ctx := context.Background() + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + mgr, mockStore, _ := newReconcileMgr(t, ctrl) + mockStore.EXPECT(). + GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1"). + Return(nil, status.Errorf(status.Internal, "store unavailable")) + + assert.Error(t, mgr.RemoveAccountGateway(ctx, "acct-1"), "a failed synthesis must fail the hook") +} + +func TestRemoveAccountGateway_NilProxyController_NoOp(t *testing.T) { + mgr := &managerImpl{reconcileCache: make(map[string]map[string]syntheticMapping)} + // Must not panic and must not query the store. + assert.NoError(t, mgr.RemoveAccountGateway(context.Background(), "acct-1")) +} + +// TestReconcile_ConcurrentWithGatewayChanges — while an account's gateway +// flaps (its policy is removed and re-added between reads), concurrent +// reconciles and RemoveAccountGateway share the cached mappings: one caches a +// mapping and sends it, another finds it gone and sends its removal. Run under +// -race: neither path may write a cached mapping, only copies of it. +func TestReconcile_ConcurrentWithGatewayChanges(t *testing.T) { + ctx := context.Background() + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl) + // gomock serialises every call on the controller's mutex, which would give + // the race detector the ordering the code under test lacks. The sends go + // through a fake that takes no lock. + mgr.proxyController = unsyncedSender{MockController: mockProxy} + provider := newReconcileTestProvider() + policy := newReconcileTestPolicy(provider.ID, "grp-eng") + + var reads atomic.Int64 + mockStore.EXPECT().GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").Return(newReconcileTestSettings(), nil).AnyTimes() + mockStore.EXPECT().GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Provider{provider}, nil).AnyTimes() + mockStore.EXPECT().GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, "acct-1"). + DoAndReturn(func(context.Context, store.LockingStrength, string) ([]*types.Policy, error) { + if reads.Add(1)%2 == 0 { + return []*types.Policy{}, nil + } + return []*types.Policy{policy}, nil + }).AnyTimes() + mockStore.EXPECT().GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Guardrail{}, nil).AnyTimes() + + var wg sync.WaitGroup + for i := 0; i < 8; i++ { + wg.Add(1) + go func(remove bool) { + defer wg.Done() + for j := 0; j < 50; j++ { + if remove && j%10 == 0 { + _ = mgr.RemoveAccountGateway(ctx, "acct-1") + continue + } + mgr.reconcile(ctx, "acct-1") + } + }(i == 0) + } + wg.Wait() +} + +// unsyncedSender answers the calls reconcile makes on every pass without any +// locking, so concurrent callers are not ordered by the fake itself. +type unsyncedSender struct { + *proxy.MockController +} + +func (unsyncedSender) GetOIDCValidationConfig() proxy.OIDCValidationConfig { + return proxy.OIDCValidationConfig{} +} + +func (unsyncedSender) SendServiceUpdateToCluster(context.Context, string, *proto.ProxyMapping, string) {} diff --git a/management/internals/modules/peers/manager.go b/management/internals/modules/peers/manager.go index 3274ec524..e944be291 100644 --- a/management/internals/modules/peers/manager.go +++ b/management/internals/modules/peers/manager.go @@ -97,7 +97,7 @@ func (m *managerImpl) GetAllPeers(ctx context.Context, accountID, userID string) return m.store.GetUserPeers(ctx, store.LockingStrengthNone, accountID, userID) } - return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") } func (m *managerImpl) GetPeerAccountID(ctx context.Context, peerID string) (string, error) { diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/manager.go b/management/internals/modules/reverseproxy/accesslogs/manager/manager.go index 4b24adaa1..d22cdb980 100644 --- a/management/internals/modules/reverseproxy/accesslogs/manager/manager.go +++ b/management/internals/modules/reverseproxy/accesslogs/manager/manager.go @@ -9,6 +9,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" + "github.com/netbirdio/netbird/management/internals/shared/db" "github.com/netbirdio/netbird/management/server/geolocation" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" @@ -18,14 +19,16 @@ import ( ) type managerImpl struct { + repo accesslogs.Repository store store.Store permissionsManager permissions.Manager geo geolocation.Geolocation cleanupCancel context.CancelFunc } -func NewManager(store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager { +func NewManager(repo accesslogs.Repository, store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager { return &managerImpl{ + repo: repo, store: store, permissionsManager: permissionsManager, geo: geo, @@ -54,7 +57,7 @@ func (m *managerImpl) SaveAccessLog(ctx context.Context, logEntry *accesslogs.Ac } } - if err := m.store.CreateAccessLog(ctx, logEntry); err != nil { + if err := m.repo.Create(ctx, logEntry); err != nil { log.WithContext(ctx).WithFields(log.Fields{ "service_id": logEntry.ServiceID, "method": logEntry.Method, @@ -82,7 +85,7 @@ func (m *managerImpl) GetAllAccessLogs(ctx context.Context, accountID, userID st log.WithContext(ctx).Warnf("failed to resolve user filters: %v", err) } - logs, totalCount, err := m.store.GetAccountAccessLogs(ctx, store.LockingStrengthNone, accountID, *filter) + logs, totalCount, err := m.repo.ListByAccount(ctx, db.LockingStrengthNone, accountID, *filter) if err != nil { return nil, 0, err } @@ -98,7 +101,7 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in } cutoffTime := time.Now().AddDate(0, 0, -retentionDays) - deletedCount, err := m.store.DeleteOldAccessLogs(ctx, cutoffTime) + deletedCount, err := m.repo.DeleteOlderThan(ctx, cutoffTime) if err != nil { log.WithContext(ctx).Errorf("failed to cleanup old access logs: %v", err) return 0, err diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go b/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go index 8e941d7e5..83dab7df6 100644 --- a/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go @@ -5,27 +5,27 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" - "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" ) func TestCleanupOldAccessLogs(t *testing.T) { tests := []struct { name string retentionDays int - setupMock func(*store.MockStore) + setupMock func(*accesslogs.MockRepository) expectedCount int64 expectedError bool }{ { name: "cleanup logs older than retention period", retentionDays: 30, - setupMock: func(mockStore *store.MockStore) { - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + setupMock: func(mockRepo *accesslogs.MockRepository) { + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) { expectedCutoff := time.Now().AddDate(0, 0, -30) timeDiff := olderThan.Sub(expectedCutoff) @@ -41,9 +41,9 @@ func TestCleanupOldAccessLogs(t *testing.T) { { name: "no logs to cleanup", retentionDays: 30, - setupMock: func(mockStore *store.MockStore) { - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + setupMock: func(mockRepo *accesslogs.MockRepository) { + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(0), nil) }, expectedCount: 0, @@ -52,8 +52,8 @@ func TestCleanupOldAccessLogs(t *testing.T) { { name: "zero retention days skips cleanup", retentionDays: 0, - setupMock: func(mockStore *store.MockStore) { - // No expectations - DeleteOldAccessLogs should not be called + setupMock: func(mockRepo *accesslogs.MockRepository) { + // No expectations - DeleteOlderThan should not be called }, expectedCount: 0, expectedError: false, @@ -61,8 +61,8 @@ func TestCleanupOldAccessLogs(t *testing.T) { { name: "negative retention days skips cleanup", retentionDays: -10, - setupMock: func(mockStore *store.MockStore) { - // No expectations - DeleteOldAccessLogs should not be called + setupMock: func(mockRepo *accesslogs.MockRepository) { + // No expectations - DeleteOlderThan should not be called }, expectedCount: 0, expectedError: false, @@ -74,11 +74,11 @@ func TestCleanupOldAccessLogs(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) - tt.setupMock(mockStore) + mockRepo := accesslogs.NewMockRepository(ctrl) + tt.setupMock(mockRepo) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx := context.Background() @@ -98,10 +98,10 @@ func TestCleanupWithExactBoundary(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) { expectedCutoff := time.Now().AddDate(0, 0, -30) timeDiff := olderThan.Sub(expectedCutoff) @@ -110,7 +110,7 @@ func TestCleanupWithExactBoundary(t *testing.T) { }) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx := context.Background() @@ -125,11 +125,11 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) // No expectations - cleanup should not run manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -139,22 +139,22 @@ func TestStartPeriodicCleanup(t *testing.T) { time.Sleep(100 * time.Millisecond) - // If DeleteOldAccessLogs was called, the test will fail due to unexpected call + // If DeleteOlderThan was called, the test will fail due to unexpected call }) t.Run("periodic cleanup runs immediately on start", func(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(2), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -171,15 +171,15 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(1), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -198,15 +198,15 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(0), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -223,15 +223,15 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(3), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -249,15 +249,15 @@ func TestStopPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(1), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx := context.Background() diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/repository.go b/management/internals/modules/reverseproxy/accesslogs/manager/repository.go new file mode 100644 index 000000000..105696ab7 --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/manager/repository.go @@ -0,0 +1,135 @@ +package manager + +import ( + "context" + "strings" + "time" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" + "github.com/netbirdio/netbird/management/internals/shared/db" + "github.com/netbirdio/netbird/shared/management/status" +) + +type sqlRepository struct { + conn *db.Conn + db *gorm.DB +} + +// NewRepository returns the access log repository backed by conn. +func NewRepository(conn *db.Conn) accesslogs.Repository { + return &sqlRepository{conn: conn, db: conn.DB(nil)} +} + +func (r *sqlRepository) WithTx(tx *db.Tx) accesslogs.Repository { + return &sqlRepository{conn: r.conn, db: r.conn.DB(tx)} +} + +func (r *sqlRepository) Create(ctx context.Context, entry *accesslogs.AccessLogEntry) error { + if err := r.db.Create(entry).Error; err != nil { + log.WithContext(ctx).WithFields(log.Fields{ + "service_id": entry.ServiceID, + "method": entry.Method, + "host": entry.Host, + "path": entry.Path, + }).Errorf("failed to create access log entry in store: %v", err) + return status.Errorf(status.Internal, "failed to create access log entry in store") + } + return nil +} + +// ListByAccount returns one page of an account's access logs together with the +// total number of entries matching the filter. +func (r *sqlRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) { + var totalCount int64 + countQuery := applyFilters(r.db.Model(&accesslogs.AccessLogEntry{}).Where("account_id = ?", accountID), filter) + if err := countQuery.Count(&totalCount).Error; err != nil { + log.WithContext(ctx).Errorf("failed to count access logs: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to count access logs") + } + + query := applyFilters(r.db.Where("account_id = ?", accountID), filter) + sortOrder := strings.ToUpper(filter.GetSortOrder()) + for _, column := range strings.Split(filter.GetSortColumn(), ",") { + if column = strings.TrimSpace(column); column != "" { + query = query.Order(column + " " + sortOrder) + } + } + query = query.Limit(filter.GetLimit()).Offset(filter.GetOffset()) + if lockStrength != db.LockingStrengthNone { + query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var logs []*accesslogs.AccessLogEntry + if err := query.Find(&logs).Error; err != nil { + log.WithContext(ctx).Errorf("failed to get access logs from store: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store") + } + + return logs, totalCount, nil +} + +func (r *sqlRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) { + result := r.db.Where("timestamp < ?", olderThan).Delete(&accesslogs.AccessLogEntry{}) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error) + return 0, status.Errorf(status.Internal, "failed to delete old access logs") + } + return result.RowsAffected, nil +} + +func applyFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB { + if filter.Search != nil { + searchPattern := "%" + *filter.Search + "%" + query = query.Where( + "id LIKE ? OR location_connection_ip LIKE ? OR host LIKE ? OR path LIKE ? OR CONCAT(host, path) LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)", + searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, + ) + } + + if filter.SourceIP != nil { + query = query.Where("location_connection_ip = ?", *filter.SourceIP) + } + + if filter.Host != nil { + query = query.Where("host = ?", *filter.Host) + } + + if filter.Path != nil { + query = query.Where("path LIKE ?", "%"+*filter.Path+"%") + } + + if filter.UserID != nil { + query = query.Where("user_id = ?", *filter.UserID) + } + + if filter.Method != nil { + query = query.Where("method = ?", *filter.Method) + } + + if filter.Status != nil { + switch *filter.Status { + case "success": + query = query.Where("(status_code >= ? AND status_code < ?)", 200, 400) + case "failed": + query = query.Where("((status_code >= ? AND status_code < ?) OR status_code >= ?)", 100, 200, 400) + } + } + + if filter.StatusCode != nil { + query = query.Where("status_code = ?", *filter.StatusCode) + } + + if filter.StartDate != nil { + query = query.Where("timestamp >= ?", *filter.StartDate) + } + + if filter.EndDate != nil { + query = query.Where("timestamp <= ?", *filter.EndDate) + } + + return query +} diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go b/management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go new file mode 100644 index 000000000..9931cc27c --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go @@ -0,0 +1,125 @@ +package manager + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" + "github.com/netbirdio/netbird/management/internals/shared/db" + "github.com/netbirdio/netbird/management/internals/shared/db/dbtest" +) + +func newTestRepository(t *testing.T) (accesslogs.Repository, *db.Conn) { + conn := dbtest.NewConn(t, &accesslogs.AccessLogEntry{}) + return NewRepository(conn), conn +} + +func newEntry(id, accountID, method string, age time.Duration) *accesslogs.AccessLogEntry { + return &accesslogs.AccessLogEntry{ + ID: id, + AccountID: accountID, + Method: method, + Host: "app.example.com", + Path: "/", + StatusCode: 200, + Timestamp: time.Now().Add(-age), + } +} + +func TestSqlRepository_ListByAccount(t *testing.T) { + repo, _ := newTestRepository(t) + ctx := context.Background() + for _, entry := range []*accesslogs.AccessLogEntry{ + newEntry("a1", "acc-a", "GET", 3*time.Hour), + newEntry("a2", "acc-a", "POST", 2*time.Hour), + newEntry("a3", "acc-a", "GET", time.Hour), + newEntry("b1", "acc-b", "GET", time.Hour), + } { + require.NoError(t, repo.Create(ctx, entry)) + } + + logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 2}) + require.NoError(t, err) + assert.EqualValues(t, 3, total) + require.Len(t, logs, 2) + assert.Equal(t, "a3", logs[0].ID) + assert.Equal(t, "a2", logs[1].ID) + + method := "GET" + logs, total, err = repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Method: &method, SortOrder: "asc"}) + require.NoError(t, err) + assert.EqualValues(t, 2, total) + require.Len(t, logs, 2) + assert.Equal(t, "a1", logs[0].ID) + assert.Equal(t, "a3", logs[1].ID) +} + +func TestSqlRepository_DeleteOlderThan(t *testing.T) { + repo, _ := newTestRepository(t) + ctx := context.Background() + require.NoError(t, repo.Create(ctx, newEntry("old", "acc", "GET", 48*time.Hour))) + require.NoError(t, repo.Create(ctx, newEntry("new", "acc", "GET", time.Hour))) + + deleted, err := repo.DeleteOlderThan(ctx, time.Now().Add(-24*time.Hour)) + require.NoError(t, err) + assert.EqualValues(t, 1, deleted) + + logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10}) + require.NoError(t, err) + assert.EqualValues(t, 1, total) + require.Len(t, logs, 1) + assert.Equal(t, "new", logs[0].ID) +} + +func TestSqlRepository_CreateInsideTransactionRollsBack(t *testing.T) { + repo, conn := newTestRepository(t) + ctx := context.Background() + failure := errors.New("abort") + + err := conn.RunInTx(ctx, func(tx *db.Tx) error { + txRepo := repo.WithTx(tx) + require.NoError(t, txRepo.Create(ctx, newEntry("tx", "acc", "GET", 0))) + _, total, err := txRepo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10}) + require.NoError(t, err) + assert.EqualValues(t, 1, total) + return failure + }) + require.ErrorIs(t, err, failure) + + _, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10}) + require.NoError(t, err) + assert.Zero(t, total) +} + +func TestSqlRepository_ListByAccount_StatusFilter(t *testing.T) { + repo, _ := newTestRepository(t) + ctx := context.Background() + statusCodes := map[string]int{"l4": 0, "info": 101, "ok": 200, "notfound": 404} + for id, code := range statusCodes { + entry := newEntry(id, "acc", "GET", time.Hour) + entry.StatusCode = code + require.NoError(t, repo.Create(ctx, entry)) + } + foreign := newEntry("foreign", "other", "GET", time.Hour) + foreign.StatusCode = 500 + require.NoError(t, repo.Create(ctx, foreign)) + + listIDs := func(status string) []string { + logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Status: &status, SortBy: "status_code", SortOrder: "asc"}) + require.NoError(t, err) + require.EqualValues(t, len(logs), total) + ids := make([]string, 0, len(logs)) + for _, entry := range logs { + ids = append(ids, entry.ID) + } + return ids + } + + assert.Equal(t, []string{"info", "notfound"}, listIDs("failed")) + assert.Equal(t, []string{"ok"}, listIDs("success")) +} diff --git a/management/internals/modules/reverseproxy/accesslogs/repository.go b/management/internals/modules/reverseproxy/accesslogs/repository.go new file mode 100644 index 000000000..5945f454c --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/repository.go @@ -0,0 +1,18 @@ +package accesslogs + +import ( + "context" + "time" + + "github.com/netbirdio/netbird/management/internals/shared/db" +) + +//go:generate go tool mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod + +// Repository persists reverse proxy access log entries. +type Repository interface { + WithTx(tx *db.Tx) Repository + Create(ctx context.Context, entry *AccessLogEntry) error + ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) + DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) +} diff --git a/management/internals/modules/reverseproxy/accesslogs/repository_mock.go b/management/internals/modules/reverseproxy/accesslogs/repository_mock.go new file mode 100644 index 000000000..7fbb05b7f --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/repository_mock.go @@ -0,0 +1,102 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./repository.go +// +// Generated by this command: +// +// mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod +// + +// Package accesslogs is a generated GoMock package. +package accesslogs + +import ( + context "context" + reflect "reflect" + time "time" + + db "github.com/netbirdio/netbird/management/internals/shared/db" + gomock "go.uber.org/mock/gomock" +) + +// MockRepository is a mock of Repository interface. +type MockRepository struct { + ctrl *gomock.Controller + recorder *MockRepositoryMockRecorder + isgomock struct{} +} + +// MockRepositoryMockRecorder is the mock recorder for MockRepository. +type MockRepositoryMockRecorder struct { + mock *MockRepository +} + +// NewMockRepository creates a new mock instance. +func NewMockRepository(ctrl *gomock.Controller) *MockRepository { + mock := &MockRepository{ctrl: ctrl} + mock.recorder = &MockRepositoryMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder { + return m.recorder +} + +// Create mocks base method. +func (m *MockRepository) Create(ctx context.Context, entry *AccessLogEntry) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Create", ctx, entry) + ret0, _ := ret[0].(error) + return ret0 +} + +// Create indicates an expected call of Create. +func (mr *MockRepositoryMockRecorder) Create(ctx, entry any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Create", reflect.TypeOf((*MockRepository)(nil).Create), ctx, entry) +} + +// DeleteOlderThan mocks base method. +func (m *MockRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteOlderThan", ctx, olderThan) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// DeleteOlderThan indicates an expected call of DeleteOlderThan. +func (mr *MockRepositoryMockRecorder) DeleteOlderThan(ctx, olderThan any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOlderThan", reflect.TypeOf((*MockRepository)(nil).DeleteOlderThan), ctx, olderThan) +} + +// ListByAccount mocks base method. +func (m *MockRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ListByAccount", ctx, lockStrength, accountID, filter) + ret0, _ := ret[0].([]*AccessLogEntry) + ret1, _ := ret[1].(int64) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// ListByAccount indicates an expected call of ListByAccount. +func (mr *MockRepositoryMockRecorder) ListByAccount(ctx, lockStrength, accountID, filter any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListByAccount", reflect.TypeOf((*MockRepository)(nil).ListByAccount), ctx, lockStrength, accountID, filter) +} + +// WithTx mocks base method. +func (m *MockRepository) WithTx(tx *db.Tx) Repository { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "WithTx", tx) + ret0, _ := ret[0].(Repository) + return ret0 +} + +// WithTx indicates an expected call of WithTx. +func (mr *MockRepositoryMockRecorder) WithTx(tx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WithTx", reflect.TypeOf((*MockRepository)(nil).WithTx), tx) +} diff --git a/management/internals/modules/reverseproxy/domain/manager/deletion_test.go b/management/internals/modules/reverseproxy/domain/manager/deletion_test.go new file mode 100644 index 000000000..b47a2bca0 --- /dev/null +++ b/management/internals/modules/reverseproxy/domain/manager/deletion_test.go @@ -0,0 +1,80 @@ +package manager + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gorilla/mux" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/activity" + nbcontext "github.com/netbirdio/netbird/management/server/context" + nbstore "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/shared/auth" +) + +func TestDeleteDomain_ServiceDependencies(t *testing.T) { + for _, tt := range []struct { + name string + domainName string + serviceHost string + accountID string + enabled bool + protected bool + }{ + {"exact", "example.com", "example.com", accountA, true, true}, + {"subdomain", "example.com", "deep.app.example.com", accountA, true, true}, + {"disabled", "example.com", "app.example.com", accountA, false, true}, + // A service is authorized by its own account's registration, so another + // account's service under this namespace is not a dependency of it. + {"other account", "example.com", "app.example.com", accountB, true, false}, + {"case and trailing dot", "example.com", "APP.EXAMPLE.COM.", accountA, true, true}, + {"suffix boundary", "example.com", "notexample.com", accountA, true, false}, + {"literal underscore", "a_b.example.com", "app.a_b.example.com", accountA, true, true}, + {"underscore wildcard", "a_b.example.com", "app.axb.example.com", accountA, true, false}, + } { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + events := captureDomainEvents(env) + d, err := env.store.CreateCustomDomain(ctx, accountA, tt.domainName, testCluster, true) + require.NoError(t, err) + svc := &rpservice.Service{ + ID: "dependent", AccountID: tt.accountID, Domain: tt.serviceHost, + Enabled: tt.enabled, ProxyCluster: testCluster, + } + require.NoError(t, env.store.CreateService(ctx, svc)) + router := mux.NewRouter() + RegisterEndpoints(router, env.manager) + deleteDomain := func() *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodDelete, "/domains/"+d.ID, nil) + req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: accountA, UserId: accountAUser}) + response := httptest.NewRecorder() + router.ServeHTTP(response, req) + return response + } + + response := deleteDomain() + if tt.protected { + require.Equal(t, http.StatusPreconditionFailed, response.Code, "dependent services must block deletion: %s", response.Body.String()) + assert.NotContains(t, response.Body.String(), tt.accountID, "the error must not reveal the service's account") + assert.NotNil(t, storedDomain(t, env.store, accountA, d.Domain), "the namespace must remain reserved") + assert.Empty(t, events.get(), "rejected deletion must not emit DomainDeleted") + stored, err := env.store.GetServiceByID(ctx, nbstore.LockingStrengthNone, tt.accountID, svc.ID) + require.NoError(t, err) + assert.Equal(t, svc.Enabled, stored.Enabled, "rejected deletion must preserve the service") + require.NoError(t, env.store.DeleteService(ctx, tt.accountID, svc.ID)) + response = deleteDomain() + } + require.Equal(t, http.StatusNoContent, response.Code, "deletion must succeed without dependencies: %s", response.Body.String()) + assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "the registration must be deleted") + captured := events.get() + require.Len(t, captured, 1, "only successful deletion may emit an event") + assert.Equal(t, activity.DomainDeleted, captured[0].Activity, "the event must describe the successful deletion") + }) + } +} diff --git a/management/internals/modules/reverseproxy/domain/manager/manager.go b/management/internals/modules/reverseproxy/domain/manager/manager.go index c0fb12e9c..1e9ecab43 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager.go @@ -357,6 +357,26 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain) } +// ValidateServiceDomain holds custom domain authorization through a service write transaction. +func (m Manager) ValidateServiceDomain(ctx context.Context, tx nbstore.Store, accountID, serviceDomain, cluster string) error { + if _, ok := ExtractClusterFromFreeDomain(serviceDomain, []string{cluster}); ok { + return nil + } + name, err := nbdomain.FromString(serviceDomain) + if err != nil { + return status.Errorf(status.InvalidArgument, "invalid service domain: %v", err) + } + customDomains, err := tx.LockCustomDomains(ctx, accountID, name) + if err != nil { + return err + } + target, match := extractClusterFromCustomDomains(serviceDomain, customDomains) + if match != customDomainValidated || target != cluster { + return status.Errorf(status.PreconditionFailed, "custom domain authorization changed; retry the service operation") + } + return nil +} + func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]string, error) { byopAddresses, err := m.proxyManager.GetActiveClusterAddressesForAccount(ctx, accountID) if err != nil { diff --git a/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go b/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go index 5c973c40e..fed402498 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go @@ -99,7 +99,7 @@ func setupDomainTest(t *testing.T) *domainTestEnv { proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) require.NoError(t, err) - _, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil) + _, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", "", nil, nil) require.NoError(t, err) resolver := &stubResolver{cnames: make(map[string]string)} diff --git a/management/internals/modules/reverseproxy/proxy/manager.go b/management/internals/modules/reverseproxy/proxy/manager.go index 26214c11b..9350ad9b9 100644 --- a/management/internals/modules/reverseproxy/proxy/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager.go @@ -11,7 +11,7 @@ import ( // Manager defines the interface for proxy operations type Manager interface { - Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) + Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) Disconnect(ctx context.Context, proxyID, sessionID string) error Heartbeat(ctx context.Context, p *Proxy) error GetActiveClusterAddresses(ctx context.Context) ([]string, error) @@ -20,6 +20,8 @@ type Manager interface { ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool + ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool CleanupStale(ctx context.Context, inactivityDuration time.Duration) error GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error) CountAccountProxies(ctx context.Context, accountID string) (int64, error) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager.go b/management/internals/modules/reverseproxy/proxy/manager/manager.go index 943766004..5a95ea94a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager.go @@ -8,6 +8,7 @@ import ( "go.opentelemetry.io/otel/metric" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" + nbversion "github.com/netbirdio/netbird/version" ) // store defines the interface for proxy persistence operations @@ -22,6 +23,8 @@ type store interface { GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool + GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error) CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error) @@ -29,6 +32,8 @@ type store interface { DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error } +const minSessionCodeVersion = "0.81.0" + // Manager handles all proxy operations type Manager struct { store store @@ -50,7 +55,7 @@ func NewManager(store store, meter metric.Meter) (*Manager, error) { // Connect registers a new proxy connection in the database. // capabilities may be nil for old proxies that do not report them. -func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) { +func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) { now := time.Now() var caps proxy.Capabilities if capabilities != nil { @@ -61,6 +66,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres SessionID: sessionID, ClusterAddress: clusterAddress, IPAddress: ipAddress, + Version: truncateVersion(version), AccountID: accountID, LastSeen: now, ConnectedAt: &now, @@ -78,6 +84,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres "sessionID": sessionID, "clusterAddress": clusterAddress, "ipAddress": ipAddress, + "version": p.Version, }).Info("proxy connected") return p, nil @@ -143,6 +150,27 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string) return m.store.GetClusterSupportsPrivate(ctx, clusterAddr) } +// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported). +func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + return m.store.GetClusterAllProxiesPrivate(ctx, clusterAddr) +} + +// ClusterSupportsSessionCode reports whether all active proxies support session codes. +func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool { + versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr) + if err != nil || len(versions) == 0 { + return false + } + + for _, version := range versions { + if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported { + return false + } + } + + return true +} + // CleanupStale removes proxies that haven't sent heartbeat in the specified duration func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error { if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil { @@ -184,3 +212,13 @@ func (m *Manager) DeleteAccountCluster(ctx context.Context, clusterAddress, acco } return nil } + +// truncateVersion cuts a proxy-reported version to the column width so an +// oversized value cannot fail the save and block the connect. +func truncateVersion(version string) string { + runes := []rune(version) + if len(runes) <= proxy.MaxVersionLength { + return version + } + return string(runes[:proxy.MaxVersionLength]) +} diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go index 5c44470a3..56806613a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go @@ -4,8 +4,10 @@ import ( "context" "errors" "fmt" + "strings" "testing" "time" + "unicode/utf8" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -20,6 +22,7 @@ type mockStore struct { updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error) getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error) + getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error) cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error) countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error) @@ -102,6 +105,15 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool { return nil } +func (m *mockStore) GetClusterAllProxiesPrivate(_ context.Context, _ string) *bool { + return nil +} +func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) { + if m.getActiveProxyVersionsFunc != nil { + return m.getActiveProxyVersionsFunc(ctx, clusterAddress) + } + return nil, nil +} func newTestManager(s store) *Manager { meter := noop.NewMeterProvider().Meter("test") @@ -112,6 +124,34 @@ func newTestManager(s store) *Manager { return m } +func TestClusterSupportsSessionCode(t *testing.T) { + tests := []struct { + name string + versions []string + storeErr error + want bool + }{ + {name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true}, + {name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}}, + {name: "missing version", versions: []string{"0.81.0", ""}}, + {name: "no active proxies"}, + {name: "store error", storeErr: errors.New("db error")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := &mockStore{ + getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) { + return tt.versions, tt.storeErr + }, + } + + got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com") + assert.Equal(t, tt.want, got) + }) + } +} + func TestConnect_WithAccountID(t *testing.T) { accountID := "acc-123" @@ -124,7 +164,7 @@ func TestConnect_WithAccountID(t *testing.T) { } mgr := newTestManager(s) - _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", &accountID, nil) + _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "0.60.0", &accountID, nil) require.NoError(t, err) require.NotNil(t, savedProxy) @@ -132,6 +172,7 @@ func TestConnect_WithAccountID(t *testing.T) { assert.Equal(t, "session-1", savedProxy.SessionID) assert.Equal(t, "cluster.example.com", savedProxy.ClusterAddress) assert.Equal(t, "10.0.0.1", savedProxy.IPAddress) + assert.Equal(t, "0.60.0", savedProxy.Version, "reported proxy version should be stored") assert.Equal(t, &accountID, savedProxy.AccountID) assert.Equal(t, proxy.StatusConnected, savedProxy.Status) assert.NotNil(t, savedProxy.ConnectedAt) @@ -147,7 +188,7 @@ func TestConnect_WithoutAccountID(t *testing.T) { } mgr := newTestManager(s) - _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", nil, nil) + _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", "", nil, nil) require.NoError(t, err) require.NotNil(t, savedProxy) @@ -155,6 +196,29 @@ func TestConnect_WithoutAccountID(t *testing.T) { assert.Equal(t, proxy.StatusConnected, savedProxy.Status) } +func TestConnect_TruncatesOversizedVersion(t *testing.T) { + var savedProxy *proxy.Proxy + s := &mockStore{ + saveProxyFunc: func(_ context.Context, p *proxy.Proxy) error { + savedProxy = p + return nil + }, + } + + // Multi-byte runes make sure the cut counts characters, as varchar does, + // and never splits a rune into invalid UTF-8. + version := strings.Repeat("ü", proxy.MaxVersionLength+10) + + mgr := newTestManager(s) + _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", version, nil, nil) + require.NoError(t, err) + + require.NotNil(t, savedProxy) + assert.Equal(t, proxy.MaxVersionLength, utf8.RuneCountInString(savedProxy.Version), "stored version should be cut to the column width") + assert.True(t, utf8.ValidString(savedProxy.Version), "stored version should remain valid UTF-8") + assert.True(t, strings.HasPrefix(version, savedProxy.Version), "stored version should be a prefix of the reported one") +} + func TestConnect_StoreError(t *testing.T) { s := &mockStore{ saveProxyFunc: func(_ context.Context, _ *proxy.Proxy) error { @@ -163,7 +227,7 @@ func TestConnect_StoreError(t *testing.T) { } mgr := newTestManager(s) - _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", nil, nil) + _, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "", nil, nil) assert.Error(t, err) } diff --git a/management/internals/modules/reverseproxy/proxy/manager_mock.go b/management/internals/modules/reverseproxy/proxy/manager_mock.go index 36d6f53fc..5f3404096 100644 --- a/management/internals/modules/reverseproxy/proxy/manager_mock.go +++ b/management/internals/modules/reverseproxy/proxy/manager_mock.go @@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration) } +// ClusterAllProxiesPrivate mocks base method. +func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr) + ret0, _ := ret[0].(*bool) + return ret0 +} + +// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate. +func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr) +} + // ClusterRequireSubdomain mocks base method. func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { m.ctrl.T.Helper() @@ -112,19 +126,33 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr) } -// Connect mocks base method. -func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) { +// ClusterSupportsSessionCode mocks base method. +func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities) + ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr) + ret0, _ := ret[0].(bool) + return ret0 +} + +// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode. +func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr) +} + +// Connect mocks base method. +func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities) ret0, _ := ret[0].(*Proxy) ret1, _ := ret[1].(error) return ret0, ret1 } // Connect indicates an expected call of Connect. -func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call { +func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities) } // CountAccountProxies mocks base method. diff --git a/management/internals/modules/reverseproxy/proxy/proxy.go b/management/internals/modules/reverseproxy/proxy/proxy.go index 4404b0d24..fdecc5bfd 100644 --- a/management/internals/modules/reverseproxy/proxy/proxy.go +++ b/management/internals/modules/reverseproxy/proxy/proxy.go @@ -9,6 +9,9 @@ const ( StatusDisconnected = "disconnected" ) +// MaxVersionLength is the width of the Version column, in characters. +const MaxVersionLength = 255 + // Capabilities describes what a proxy can handle, as reported via gRPC. // Nil fields mean the proxy never reported this capability. type Capabilities struct { @@ -31,6 +34,7 @@ type Proxy struct { SessionID string `gorm:"type:varchar(36)"` ClusterAddress string `gorm:"type:varchar(255);not null;index:idx_proxy_cluster_status"` IPAddress string `gorm:"type:varchar(45)"` + Version string `gorm:"type:varchar(255)"` AccountID *string `gorm:"type:varchar(255);index:idx_proxy_account_id"` LastSeen time.Time `gorm:"not null;index:idx_proxy_last_seen"` ConnectedAt *time.Time diff --git a/management/internals/modules/reverseproxy/proxytoken/handler.go b/management/internals/modules/reverseproxy/proxytoken/handler.go index ed098a6dd..d8578db4c 100644 --- a/management/internals/modules/reverseproxy/proxytoken/handler.go +++ b/management/internals/modules/reverseproxy/proxytoken/handler.go @@ -1,6 +1,7 @@ package proxytoken import ( + "context" "encoding/json" "net/http" "time" @@ -18,13 +19,29 @@ import ( "github.com/netbirdio/netbird/shared/management/status" ) +// RevocationGuard vetoes the tenant-facing revocation of a proxy access +// token. Implementations are supplied by integrations; none is installed by +// default, so every token the caller's account owns may be revoked. It is +// consulted after the ownership check and before the token is revoked. A +// returned status error is written with util.WriteError: its type selects the +// HTTP status and its message is shown to the caller, so it must not carry +// internal detail. Any other error is reported as a generic internal error. +type RevocationGuard interface { + CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error +} + type handler struct { store store.Store permissionsManager permissions.Manager + // revocationGuard vetoes revocations. Optional — when nil every owned + // token may be revoked. + revocationGuard RevocationGuard } -func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, router *mux.Router) { - h := &handler{store: s, permissionsManager: permissionsManager} +// RegisterEndpoints registers the proxy token endpoints. revocationGuard is +// optional; pass nil for no revocation policy. +func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, revocationGuard RevocationGuard, router *mux.Router) { + h := &handler{store: s, permissionsManager: permissionsManager, revocationGuard: revocationGuard} router.HandleFunc("/reverse-proxies/proxy-tokens", h.listTokens).Methods("GET", "OPTIONS") router.HandleFunc("/reverse-proxies/proxy-tokens", h.createToken).Methods("POST", "OPTIONS") router.HandleFunc("/reverse-proxies/proxy-tokens/{tokenId}", h.revokeToken).Methods("DELETE", "OPTIONS") @@ -154,6 +171,13 @@ func (h *handler) revokeToken(w http.ResponseWriter, r *http.Request) { return } + if h.revocationGuard != nil { + if err := h.revocationGuard.CheckProxyAccessTokenRevocation(ctx, token); err != nil { + util.WriteError(ctx, err, w) + return + } + } + if err := h.store.RevokeProxyAccessToken(ctx, tokenID); err != nil { util.WriteErrorResponse("failed to revoke token", http.StatusInternalServerError, w) return diff --git a/management/internals/modules/reverseproxy/proxytoken/handler_test.go b/management/internals/modules/reverseproxy/proxytoken/handler_test.go index c71fe59f6..da37261c4 100644 --- a/management/internals/modules/reverseproxy/proxytoken/handler_test.go +++ b/management/internals/modules/reverseproxy/proxytoken/handler_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "net/http" "net/http/httptest" "testing" @@ -22,6 +23,7 @@ import ( "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/auth" "github.com/netbirdio/netbird/shared/management/http/api" + "github.com/netbirdio/netbird/shared/management/status" ) func authContext(accountID, userID string) context.Context { @@ -273,3 +275,152 @@ func TestRevokeToken_ManagementWideToken(t *testing.T) { h.revokeToken(w, req) assert.Equal(t, http.StatusNotFound, w.Code) } + +type revocationGuardFunc func(ctx context.Context, token *types.ProxyAccessToken) error + +func (f revocationGuardFunc) CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error { + return f(ctx, token) +} + +func TestRevokeToken_GuardRefuses(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + accountID := "acc-123" + + // No RevokeProxyAccessToken expectation: a refused revocation must not + // reach the store. + mockStore := store.NewMockStore(ctrl) + mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{ + ID: "tok-1", + AccountID: &accountID, + }, nil) + + permsMgr := permissions.NewMockManager(ctrl) + permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil) + + var checked *types.ProxyAccessToken + h := &handler{ + store: mockStore, + permissionsManager: permsMgr, + revocationGuard: revocationGuardFunc(func(_ context.Context, token *types.ProxyAccessToken) error { + checked = token + return status.Errorf(status.PreconditionFailed, "token is in use") + }), + } + + req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil) + req = req.WithContext(authContext(accountID, "user-1")) + req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"}) + w := httptest.NewRecorder() + + h.revokeToken(w, req) + assert.Equal(t, http.StatusPreconditionFailed, w.Code) + assert.Contains(t, w.Body.String(), "token is in use") + require.NotNil(t, checked) + assert.Equal(t, "tok-1", checked.ID) +} + +func TestRevokeToken_GuardAllows(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + accountID := "acc-123" + + mockStore := store.NewMockStore(ctrl) + mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{ + ID: "tok-1", + AccountID: &accountID, + }, nil) + mockStore.EXPECT().RevokeProxyAccessToken(gomock.Any(), "tok-1").Return(nil) + + permsMgr := permissions.NewMockManager(ctrl) + permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil) + + h := &handler{ + store: mockStore, + permissionsManager: permsMgr, + revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error { + return nil + }), + } + + req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil) + req = req.WithContext(authContext(accountID, "user-1")) + req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"}) + w := httptest.NewRecorder() + + h.revokeToken(w, req) + assert.Equal(t, http.StatusOK, w.Code) +} + +func TestRevokeToken_GuardFailure(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + accountID := "acc-123" + + // No RevokeProxyAccessToken expectation: a guard that cannot decide must + // not let the revocation through. + mockStore := store.NewMockStore(ctrl) + mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{ + ID: "tok-1", + AccountID: &accountID, + }, nil) + + permsMgr := permissions.NewMockManager(ctrl) + permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil) + + h := &handler{ + store: mockStore, + permissionsManager: permsMgr, + revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error { + return errors.New("connection refused") + }), + } + + req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil) + req = req.WithContext(authContext(accountID, "user-1")) + req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"}) + w := httptest.NewRecorder() + + h.revokeToken(w, req) + assert.Equal(t, http.StatusInternalServerError, w.Code) + assert.Contains(t, w.Body.String(), "internal server error") + assert.NotContains(t, w.Body.String(), "connection refused") +} + +func TestRevokeToken_GuardNotConsultedForForeignToken(t *testing.T) { + ctrl := gomock.NewController(t) + defer ctrl.Finish() + + otherAccount := "acc-other" + + mockStore := store.NewMockStore(ctrl) + mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{ + ID: "tok-1", + AccountID: &otherAccount, + }, nil) + + permsMgr := permissions.NewMockManager(ctrl) + permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), "acc-123", "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil) + + // A foreign token must read as not found, not reveal through the guard's + // answer that it belongs to some account's managed proxy. + h := &handler{ + store: mockStore, + permissionsManager: permsMgr, + revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error { + t.Fatal("guard consulted for a token the caller does not own") + return nil + }), + } + + req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil) + req = req.WithContext(authContext("acc-123", "user-1")) + req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"}) + w := httptest.NewRecorder() + + h.revokeToken(w, req) + assert.Equal(t, http.StatusNotFound, w.Code) +} diff --git a/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go b/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go index ccb955cd8..f491d01c3 100644 --- a/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go +++ b/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go @@ -30,7 +30,7 @@ func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) { proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) require.NoError(t, err) - _, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil) + _, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", "", nil, nil) require.NoError(t, err) accountMgr := &mock_server.MockAccountManager{ @@ -125,3 +125,53 @@ func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) { require.NoError(t, err) assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain") } + +func TestCreateService_DomainDeletedBeforeWrite(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupIntegrationTest(t) + withRealDomainManager(t, mgr, testStore) + + d, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true) + require.NoError(t, err) + svc := newTestService("app.proven.example.com") + require.NoError(t, mgr.initializeServiceForCreate(ctx, testAccountID, svc)) + + // Delete after the initial authorization check, before the service transaction starts. + require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID)) + err = mgr.persistNewService(ctx, testAccountID, svc) + require.Error(t, err, "an earlier validation result must not authorize a deleted registration") + sErr, ok := status.FromError(err) + require.True(t, ok, "the caller must receive a typed precondition error") + assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the service must require current domain authorization") + services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, err) + assert.Empty(t, services, "the failed write must not leave a service") +} + +func TestUpdateService_DomainDeletedBeforeWrite(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupIntegrationTest(t) + withRealDomainManager(t, mgr, testStore) + _, err := testStore.CreateCustomDomain(ctx, testAccountID, "original.example.com", validationTestCluster, true) + require.NoError(t, err) + d, err := testStore.CreateCustomDomain(ctx, testAccountID, "destination.example.com", validationTestCluster, true) + require.NoError(t, err) + svc, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.original.example.com")) + require.NoError(t, err) + moved := svc.Copy() + moved.Domain = "app.destination.example.com" + cluster, err := mgr.resolveEffectiveCluster(ctx, testAccountID, moved) + require.NoError(t, err) + + require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID)) + err = testStore.ExecuteInTransaction(ctx, func(tx store.Store) error { + return mgr.executeServiceUpdate(ctx, tx, testAccountID, moved, &serviceUpdateInfo{}, nil, cluster) + }) + require.Error(t, err, "a domain deleted after cluster resolution must reject the update") + sErr, ok := status.FromError(err) + require.True(t, ok, "the caller must receive a typed precondition error") + assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the move must require current domain authorization") + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, svc.ID) + require.NoError(t, err) + assert.Equal(t, svc.Domain, stored.Domain, "the service must retain its authorized domain") +} diff --git a/management/internals/modules/reverseproxy/service/manager/manager.go b/management/internals/modules/reverseproxy/service/manager/manager.go index 9c7f95eb4..900b7759f 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager.go +++ b/management/internals/modules/reverseproxy/service/manager/manager.go @@ -74,6 +74,7 @@ const unknownHostPlaceholder = "unknown" // ClusterDeriver derives the proxy cluster from a domain. type ClusterDeriver interface { DeriveClusterFromDomain(ctx context.Context, accountID, domain string) (string, error) + ValidateServiceDomain(ctx context.Context, tx store.Store, accountID, domain, cluster string) error GetClusterDomains() []string } @@ -83,6 +84,7 @@ type CapabilityProvider interface { ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool } type Manager struct { @@ -331,7 +333,14 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc * return err } + if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil { + return err + } + return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil { + return err + } if svc.Domain != "" { if err := m.checkDomainAvailable(ctx, transaction, svc.Domain, ""); err != nil { return err @@ -365,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service) return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster) } +// validatePrivateClusterTargets rejects cluster and direct upstream targets unless +// every active proxy in the service's cluster reports the private capability. The +// mapping reaches all proxies in the cluster, so one non-private proxy would serve +// these targets too. An unreported capability is treated as unsupported. Must be +// called outside a transaction, like clusterCustomPorts. +func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error { + target := firstPrivateClusterTarget(targets) + if target == nil { + return nil + } + + if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private { + return nil + } + + if target.TargetType == service.TargetTypeCluster { + return status.Errorf(status.InvalidArgument, + "target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it", + service.TargetTypeCluster, cluster) + } + return status.Errorf(status.InvalidArgument, + "direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster) +} + +// firstPrivateClusterTarget returns the first target that only a private cluster may serve. +func firstPrivateClusterTarget(targets []*service.Target) *service.Target { + for _, target := range targets { + if target == nil { + continue + } + if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream { + return target + } + } + return nil +} + // ensureL4Port auto-assigns a listen port when needed and validates cluster support. // customPorts must be pre-computed via clusterCustomPorts before entering a transaction. func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error { @@ -460,7 +506,14 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee return err } + if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil { + return err + } + return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil { + return err + } if err := m.validateEphemeralPreconditions(ctx, transaction, accountID, peerID, svc); err != nil { return err } @@ -577,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se return nil, err } + if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil { + return nil, err + } + // Validate subdomain requirement *before* the transaction: the underlying // capability lookup talks to the main DB pool, and SQLite's single-connection // pool would self-deadlock if this ran while the tx already held the only @@ -622,6 +679,9 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string, } func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error { + if err := m.validateServiceDomain(ctx, transaction, accountID, service, effectiveCluster); err != nil { + return err + } existingService, err := transaction.GetServiceByID(ctx, store.LockingStrengthUpdate, accountID, service.ID) if err != nil { return err @@ -677,6 +737,13 @@ func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.St return nil } +func (m *Manager) validateServiceDomain(ctx context.Context, tx store.Store, accountID string, svc *service.Service, cluster string) error { + if m.clusterDeriver == nil { + return nil + } + return m.clusterDeriver.ValidateServiceDomain(ctx, tx, accountID, svc.Domain, cluster) +} + // validateL4PortDiffOnClusterDiff checks if custom L4 ports are configured and validates port changes across clusters. // It ensures no port changes if custom ports are unsupported for a given cluster and protocol mode. // Returns an error if validation fails, otherwise returns nil. diff --git a/management/internals/modules/reverseproxy/service/manager/manager_test.go b/management/internals/modules/reverseproxy/service/manager/manager_test.go index dd0edec60..2ec1af1a2 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/service/manager/manager_test.go @@ -433,8 +433,8 @@ func TestDeletePeerService_SourcePeerValidation(t *testing.T) { newProxyServer := func(t *testing.T) *nbgrpc.ProxyServiceServer { t.Helper() tokenStore := nbgrpc.NewOneTimeTokenStore(context.Background(), testCacheStore(t)) - pkceStore := nbgrpc.NewPKCEVerifierStore(context.Background(), testCacheStore(t)) - srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) + singleUseStore := nbgrpc.NewSingleUseStore(context.Background(), testCacheStore(t)) + srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) return srv } @@ -655,6 +655,10 @@ func (d *testClusterDeriver) GetClusterDomains() []string { return d.domains } +func (d *testClusterDeriver) ValidateServiceDomain(context.Context, store.Store, string, string, string) error { + return nil +} + const ( testAccountID = "test-account" testPeerID = "test-peer-1" @@ -722,8 +726,8 @@ func setupIntegrationTest(t *testing.T) (*Manager, store.Store) { } tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t)) - pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t)) - proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t)) + proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter("")) require.NoError(t, err) @@ -1146,8 +1150,8 @@ func TestDeleteService_DeletesTargets(t *testing.T) { mockAcct := account.NewMockManager(ctrl) tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t)) - pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t)) - proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t)) + proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter("")) require.NoError(t, err) diff --git a/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go new file mode 100644 index 000000000..1f507294e --- /dev/null +++ b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go @@ -0,0 +1,218 @@ +package manager + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" + "go.uber.org/mock/gomock" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" + proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/shared/management/status" +) + +// setupPrivateClusterTest wires the real proxy manager as the capability +// provider and connects one proxy to testCluster reporting the given private +// capability. A nil private connects no proxy, so the capability is unreported. +func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) { + t.Helper() + + mgr, testStore := setupIntegrationTest(t) + + proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) + require.NoError(t, err) + mgr.capabilities = proxyMgr + + if private != nil { + connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private}) + } + + return mgr, testStore +} + +func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) { + t.Helper() + _, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps) + require.NoError(t, err) +} + +func clusterTarget() *rpservice.Target { + return &rpservice.Target{ + TargetId: testCluster, + TargetType: rpservice.TargetTypeCluster, + Host: "backend.lan", + Port: 8080, + Protocol: "http", + Enabled: true, + Options: rpservice.TargetOptions{DirectUpstream: true}, + } +} + +func directUpstreamPeerTarget() *rpservice.Target { + return &rpservice.Target{ + TargetId: testPeerID, + TargetType: rpservice.TargetTypePeer, + Host: "backend.lan", + Port: 8080, + Protocol: "http", + Enabled: true, + Options: rpservice.TargetOptions{DirectUpstream: true}, + } +} + +func TestCreateService_PrivateClusterTargets(t *testing.T) { + tests := []struct { + name string + private *bool + target *rpservice.Target + wantErr string + }{ + {name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()}, + {name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()}, + {name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + {name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, tc.private) + + svc := newTestService("app.test.netbird.io") + svc.Targets = []*rpservice.Target{tc.target} + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, svc) + + services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, listErr) + + if tc.wantErr == "" { + require.NoError(t, err) + assert.Len(t, services, 1, "the service should be persisted") + return + } + + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + sErr, ok := status.FromError(err) + require.True(t, ok, "the caller must receive a typed error") + assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument") + assert.Empty(t, services, "a rejected service must not be persisted") + }) + } +} + +// A cluster where only some proxies run in private mode must not accept these +// targets: the mapping is delivered to every proxy in the cluster, so the +// non-private ones would serve the target from their host network as well. +func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) { + tests := []struct { + name string + secondCaps *proxy.Capabilities + }{ + {name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}}, + {name: "second proxy predates capability reporting", secondCaps: nil}, + } + + for _, tc := range tests { + for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} { + t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(true)) + connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps) + + svc := newTestService("app.test.netbird.io") + svc.Targets = []*rpservice.Target{target} + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, svc) + require.Error(t, err, "a cluster with a non-private proxy must not accept the target") + assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled") + + services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, err) + assert.Empty(t, services, "a rejected service must not be persisted") + }) + } + } +} + +func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) { + ctx := context.Background() + mgr, _ := setupPrivateClusterTest(t, boolPtr(false)) + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err, "a peer target without direct upstream must not need a private cluster") +} + +func TestUpdateService_PrivateClusterTargets(t *testing.T) { + tests := []struct { + name string + target *rpservice.Target + wantErr string + }{ + {name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(false)) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err) + + updated := newTestService("app.test.netbird.io") + updated.ID = created.ID + updated.AccountID = testAccountID + updated.Targets = []*rpservice.Target{tc.target} + + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + require.Len(t, stored.Targets, 1) + assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged") + assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled") + }) + } +} + +func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(true)) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err) + + updated := newTestService("app.test.netbird.io") + updated.ID = created.ID + updated.AccountID = testAccountID + updated.Targets = []*rpservice.Target{clusterTarget()} + + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated) + require.NoError(t, err) + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + require.Len(t, stored.Targets, 1) + assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored") +} + +func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) { + ctrl := gomock.NewController(t) + // No ClusterAllProxiesPrivate expectation: a lookup would fail the test. + mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)} + + targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}} + require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster)) +} diff --git a/management/internals/modules/zones/manager/manager.go b/management/internals/modules/zones/manager/manager.go index d5348d3d0..6f6ba6c40 100644 --- a/management/internals/modules/zones/manager/manager.go +++ b/management/internals/modules/zones/manager/manager.go @@ -3,15 +3,16 @@ package manager import ( "context" "fmt" + "slices" "github.com/netbirdio/netbird/management/internals/modules/zones" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, } zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups) + var snap *affectedpeers.Snapshot + change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain) if err != nil { @@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, } if err = transaction.CreateZone(ctx, zone); err != nil { - return fmt.Errorf("failed to create zone: %w", err) + return fmt.Errorf("create zone: %w", err) + } + + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { + return fmt.Errorf("increment network serial: %w", err) } return nil @@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta()) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) + return zone, nil } @@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, return nil, status.NewPermissionDeniedError() } - zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID) - if err != nil { - return nil, fmt.Errorf("failed to get zone: %w", err) - } - - if zone.Domain != updatedZone.Domain { - return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated") - } - - zone.Name = updatedZone.Name - zone.Enabled = updatedZone.Enabled - zone.EnableSearchDomain = updatedZone.EnableSearchDomain - zone.DistributionGroups = updatedZone.DistributionGroups + var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID) + if err != nil { + return fmt.Errorf("get zone: %w", err) + } + + if zone.Domain != updatedZone.Domain { + return status.Errorf(status.InvalidArgument, "zone domain cannot be updated") + } + + oldGroups := zone.DistributionGroups + zone.Name = updatedZone.Name + zone.Enabled = updatedZone.Enabled + zone.EnableSearchDomain = updatedZone.EnableSearchDomain + zone.DistributionGroups = updatedZone.DistributionGroups + for _, groupID := range zone.DistributionGroups { _, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID) if err != nil { @@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, } if err = transaction.UpdateZone(ctx, zone); err != nil { - return fmt.Errorf("failed to update zone: %w", err) + return fmt.Errorf("update zone: %w", err) + } + + change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { + return fmt.Errorf("increment network serial: %w", err) } return nil @@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta()) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return zone, nil } @@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID return status.NewPermissionDeniedError() } - zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) - if err != nil { - return fmt.Errorf("failed to get zone: %w", err) - } - + var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change var eventsToStore []func() + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) + if err != nil { + return fmt.Errorf("get zone: %w", err) + } + + // Load before delete: the post-delete state no longer references the groups. + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID) if err != nil { return fmt.Errorf("failed to get records: %w", err) @@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID event() } - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return nil } diff --git a/management/internals/modules/zones/records/manager/manager.go b/management/internals/modules/zones/records/manager/manager.go index b041aca30..16839c1b4 100644 --- a/management/internals/modules/zones/records/manager/manager.go +++ b/management/internals/modules/zones/records/manager/manager.go @@ -9,11 +9,11 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/zones/records" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI } var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL) err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to create dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return record, nil } @@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI var zone *zones.Zone var record *records.Record + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) @@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to update dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return record, nil } @@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI var record *records.Record var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) @@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to delete dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return nil } diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index ea999d82b..34c8363f1 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -32,17 +32,18 @@ import ( networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" + "github.com/netbirdio/netbird/management/internals/shared/db" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/activity" activitystore "github.com/netbirdio/netbird/management/server/activity/store" nbcache "github.com/netbirdio/netbird/management/server/cache" nbContext "github.com/netbirdio/netbird/management/server/context" nbhttp "github.com/netbirdio/netbird/management/server/http" - "github.com/netbirdio/netbird/management/server/http/middleware" "github.com/netbirdio/netbird/management/server/idp" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/telemetry" mgmtProto "github.com/netbirdio/netbird/shared/management/proto" + "github.com/netbirdio/netbird/shared/ratelimit" "github.com/netbirdio/netbird/util/crypt" ) @@ -84,9 +85,20 @@ func (s *BaseServer) CacheStore() nbcache.Store { }) } +// DBConn opens the database connection shared by the store and the domain repositories. +func (s *BaseServer) DBConn() *db.Conn { + return Create(s, func() *db.Conn { + conn, err := store.OpenConn(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir) + if err != nil { + log.Fatalf("failed to open database connection: %v", err) + } + return conn + }) +} + func (s *BaseServer) Store() store.Store { return Create(s, func() store.Store { - store, err := store.NewStore(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir, s.Metrics(), false) + store, err := store.NewSqlStore(context.Background(), s.DBConn(), s.Metrics(), false) if err != nil { log.Fatalf("failed to create store: %v", err) } @@ -147,7 +159,7 @@ func (s *BaseServer) EventStore() activity.Store { func (s *BaseServer) APIHandler() http.Handler { return Create(s, func() http.Handler { - httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager()) + httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager(), nil) if err != nil { log.Fatalf("failed to create API handler: %v", err) } @@ -171,10 +183,10 @@ func (s *BaseServer) Router() *mux.Router { }) } -func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter { - return Create(s, func() *middleware.APIRateLimiter { - cfg, enabled := middleware.RateLimiterConfigFromEnv() - limiter := middleware.NewAPIRateLimiter(cfg) +func (s *BaseServer) RateLimiter() *ratelimit.APIRateLimiter { + return Create(s, func() *ratelimit.APIRateLimiter { + cfg, enabled := ratelimit.RateLimiterConfigFromEnv() + limiter := ratelimit.NewAPIRateLimiter(cfg) limiter.SetEnabled(enabled) return limiter }) @@ -236,7 +248,7 @@ func (s *BaseServer) GRPCServer() *grpc.Server { func (s *BaseServer) ReverseProxyGRPCServer() *nbgrpc.ProxyServiceServer { return Create(s, func() *nbgrpc.ProxyServiceServer { - proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.PKCEVerifierStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store()) + proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.SingleUseStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store()) s.AfterInit(func(s *BaseServer) { proxyService.SetServiceManager(s.ServiceManager()) proxyService.SetActivityManager(s.ProxyActivityManager()) @@ -293,9 +305,9 @@ func (s *BaseServer) ProxyTokenStore() *nbgrpc.OneTimeTokenStore { }) } -func (s *BaseServer) PKCEVerifierStore() *nbgrpc.PKCEVerifierStore { - return Create(s, func() *nbgrpc.PKCEVerifierStore { - return nbgrpc.NewPKCEVerifierStore(context.Background(), s.CacheStore()) +func (s *BaseServer) SingleUseStore() *nbgrpc.SingleUseStore { + return Create(s, func() *nbgrpc.SingleUseStore { + return nbgrpc.NewSingleUseStore(context.Background(), s.CacheStore()) }) } @@ -308,7 +320,7 @@ func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager { func (s *BaseServer) AccessLogsManager() accesslogs.Manager { return Create(s, func() accesslogs.Manager { - accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager()) + accessLogManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(s.DBConn()), s.Store(), s.PermissionsManager(), s.GeoLocationManager()) accessLogManager.StartPeriodicCleanup( context.Background(), s.Config.ReverseProxy.AccessLogRetentionDays, diff --git a/management/internals/server/modules.go b/management/internals/server/modules.go index 6b1365f3b..4840e40ad 100644 --- a/management/internals/server/modules.go +++ b/management/internals/server/modules.go @@ -103,6 +103,7 @@ func (s *BaseServer) AccountManager() account.Manager { s.AfterInit(func(s *BaseServer) { accountManager.SetServiceManager(s.ServiceManager()) + accountManager.AddAccountDeletionHook(s.AgentNetworkManager().RemoveAccountGateway) }) return accountManager diff --git a/management/internals/server/server.go b/management/internals/server/server.go index a1b58fdf1..6d51745a7 100644 --- a/management/internals/server/server.go +++ b/management/internals/server/server.go @@ -23,6 +23,8 @@ import ( "github.com/netbirdio/netbird/management/server/idp" "github.com/netbirdio/netbird/management/server/metrics" "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/shared/lifecycle" + "github.com/netbirdio/netbird/shared/profiling" "github.com/netbirdio/netbird/util/wsproxy" wsproxyserver "github.com/netbirdio/netbird/util/wsproxy/server" "github.com/netbirdio/netbird/version" @@ -36,6 +38,8 @@ const ( DefaultSelfHostedDomain = "netbird.selfhosted" ContainerKeyBaseServer = "baseServer" + + applicationName = "management" ) type Server interface { @@ -82,6 +86,8 @@ type BaseServer struct { errCh chan error wg sync.WaitGroup cancel context.CancelFunc + + lifecycle.StopHandlers } // Config holds the configuration parameters for creating a new server @@ -117,6 +123,9 @@ func NewServer(cfg *Config) *BaseServer { } s.container[ContainerKeyBaseServer] = s + stopProfiling := profiling.Start(applicationName) + s.OnStop(stopProfiling) + return s } @@ -126,6 +135,14 @@ func (s *BaseServer) AfterInit(fn func(s *BaseServer)) { // Start begins listening for HTTP requests on the configured address func (s *BaseServer) Start(ctx context.Context) error { + if err := s.start(ctx); err != nil { + s.RunStopHandlers() + return err + } + return nil +} + +func (s *BaseServer) start(ctx context.Context) error { srvCtx, cancel := context.WithCancel(ctx) s.cancel = cancel s.errCh = make(chan error, 4) @@ -278,6 +295,7 @@ func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) { func (s *BaseServer) Stop() error { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() + defer s.RunStopHandlers() if s.domainCleanupStop != nil { s.domainCleanupStop() } diff --git a/management/internals/shared/db/conn.go b/management/internals/shared/db/conn.go new file mode 100644 index 000000000..8edaa1e4e --- /dev/null +++ b/management/internals/shared/db/conn.go @@ -0,0 +1,121 @@ +package db + +import ( + "context" + "fmt" + "os" + "runtime" + "strconv" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" +) + +const ( + defaultTransactionTimeout = 5 * time.Minute + connMaxLifetime = time.Hour + connMaxIdleTime = 3 * time.Minute +) + +// TxMetrics receives the duration of every committed top-level transaction. +type TxMetrics interface { + CountTransactionDuration(duration time.Duration) +} + +// Conn is the database connection shared by all repositories: one gorm handle, +// the pgx pool of a Postgres deployment and the engine they talk to. +type Conn struct { + db *gorm.DB + pool *pgxpool.Pool + engine Engine + txTimeout time.Duration + metrics TxMetrics +} + +// NewConn takes ownership of an open gorm handle and pool once it returns +// without error, applying the connection limits and transaction timeout +// configured through the environment. +func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) { + sqlDB, err := gormDB.DB() + if err != nil { + return nil, err + } + + txTimeout := defaultTransactionTimeout + if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" { + if parsed, err := time.ParseDuration(v); err == nil { + txTimeout = parsed + } + } + log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout) + + conns := runtime.NumCPU() + configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS")) + connsConfigured := err == nil + if connsConfigured { + conns = configuredConns + } + if engine == SqliteStoreEngine { + if connsConfigured { + log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1") + } + conns = 1 + } + + sqlDB.SetMaxOpenConns(conns) + sqlDB.SetMaxIdleConns(conns) + sqlDB.SetConnMaxLifetime(connMaxLifetime) + sqlDB.SetConnMaxIdleTime(connMaxIdleTime) + + log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v", + conns, conns, connMaxLifetime, connMaxIdleTime) + + return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil +} + +// DB returns the handle a query must run on: the transaction when tx is set, +// otherwise the shared connection. +func (c *Conn) DB(tx *Tx) *gorm.DB { + if tx != nil { + return tx.db + } + return c.db +} + +// Pool returns the pgx pool for read paths that bypass gorm. It is nil on +// engines other than Postgres and inside a transaction, where the pool would +// not see the uncommitted writes. +func (c *Conn) Pool(tx *Tx) *pgxpool.Pool { + if tx != nil { + return nil + } + return c.pool +} + +func (c *Conn) Engine() Engine { + return c.engine +} + +// SetTxMetrics registers the sink that receives transaction durations. +func (c *Conn) SetTxMetrics(metrics TxMetrics) { + c.metrics = metrics +} + +// AutoMigrate creates or updates the tables of the given models. +func (c *Conn) AutoMigrate(models ...any) error { + return c.db.AutoMigrate(models...) +} + +// Close releases the gorm connection and the pgx pool. +func (c *Conn) Close() error { + if c.pool != nil { + c.pool.Close() + } + sqlDB, err := c.db.DB() + if err != nil { + return fmt.Errorf("get db: %w", err) + } + return sqlDB.Close() +} diff --git a/management/internals/shared/db/conn_test.go b/management/internals/shared/db/conn_test.go new file mode 100644 index 000000000..4f1a234b2 --- /dev/null +++ b/management/internals/shared/db/conn_test.go @@ -0,0 +1,148 @@ +package db + +import ( + "context" + "errors" + "path/filepath" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +type testRow struct { + ID uint `gorm:"primaryKey"` + Name string +} + +func openTestConn(t *testing.T) *Conn { + t.Helper() + conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, conn.Close()) }) + require.NoError(t, conn.AutoMigrate(&testRow{})) + return conn +} + +func countRows(t *testing.T, conn *Conn) int64 { + t.Helper() + var count int64 + require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error) + return count +} + +func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) { + t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s") + conn := openTestConn(t) + assert.Equal(t, time.Second, conn.txTimeout) + assert.Equal(t, SqliteStoreEngine, conn.Engine()) +} + +func TestRunInTx_CommitsOnSuccess(t *testing.T) { + conn := openTestConn(t) + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + return conn.DB(tx).Create(&testRow{Name: "a"}).Error + }) + require.NoError(t, err) + assert.EqualValues(t, 1, countRows(t, conn)) +} + +func TestRunInTx_RollsBackOnError(t *testing.T) { + conn := openTestConn(t) + failure := errors.New("boom") + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error) + return failure + }) + require.ErrorIs(t, err, failure) + assert.EqualValues(t, 0, countRows(t, conn)) +} + +func TestRunInTx_RollsBackOnPanic(t *testing.T) { + conn := openTestConn(t) + + require.Panics(t, func() { + _ = conn.RunInTx(context.Background(), func(tx *Tx) error { + require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error) + panic("boom") + }) + }) + assert.EqualValues(t, 0, countRows(t, conn)) +} + +func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) { + t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms") + conn := openTestConn(t) + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + time.Sleep(100 * time.Millisecond) + return conn.DB(tx).Create(&testRow{Name: "a"}).Error + }) + require.ErrorIs(t, err, context.DeadlineExceeded) + assert.EqualValues(t, 0, countRows(t, conn)) +} + +func TestRunInTx_ReportsDurationToMetrics(t *testing.T) { + conn := openTestConn(t) + metrics := &recordingMetrics{} + conn.SetTxMetrics(metrics) + + require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil })) + assert.Equal(t, 1, metrics.calls) +} + +func TestConn_DBSelectsTransactionHandle(t *testing.T) { + conn := openTestConn(t) + assert.Same(t, conn.db, conn.DB(nil)) + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + assert.Same(t, tx.db, conn.DB(tx)) + assert.NotSame(t, conn.db, conn.DB(tx)) + return nil + }) + require.NoError(t, err) +} + +func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) { + conn := openTestConn(t) + conn.pool = &pgxpool.Pool{} + defer func() { conn.pool = nil }() + + assert.Same(t, conn.pool, conn.Pool(nil)) + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + assert.Nil(t, conn.Pool(tx)) + return nil + }) + require.NoError(t, err) +} + +type recordingMetrics struct { + calls int +} + +func (m *recordingMetrics) CountTransactionDuration(time.Duration) { + m.calls++ +} + +func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) { + t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7") + + gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig()) + require.NoError(t, err) + conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + sqlDB, err := conn.DB(nil).DB() + require.NoError(t, err) + assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections) + + sqliteDB, err := openTestConn(t).DB(nil).DB() + require.NoError(t, err) + assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections) +} diff --git a/management/internals/shared/db/dbtest/dbtest.go b/management/internals/shared/db/dbtest/dbtest.go new file mode 100644 index 000000000..0f0fe0356 --- /dev/null +++ b/management/internals/shared/db/dbtest/dbtest.go @@ -0,0 +1,23 @@ +package dbtest + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/shared/db" +) + +// NewConn opens a fresh SQLite database in a temporary directory, migrates the +// given models and closes the connection when the test ends. It ignores +// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never +// touched, and is safe to call from parallel tests. +func NewConn(t testing.TB, models ...any) *db.Conn { + t.Helper() + conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + require.NoError(t, conn.AutoMigrate(models...)) + return conn +} diff --git a/management/internals/shared/db/dbtest/dbtest_test.go b/management/internals/shared/db/dbtest/dbtest_test.go new file mode 100644 index 000000000..3f8566876 --- /dev/null +++ b/management/internals/shared/db/dbtest/dbtest_test.go @@ -0,0 +1,31 @@ +package dbtest + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/shared/db" +) + +func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) { + override := filepath.Join(t.TempDir(), "configured.db") + t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override) + + conn := NewConn(t) + + assert.Equal(t, db.SqliteStoreEngine, conn.Engine()) + _, err := os.Stat(override) + require.ErrorIs(t, err, os.ErrNotExist) +} + +func TestNewConn_Parallel(t *testing.T) { + t.Parallel() + + conn := NewConn(t) + + assert.Equal(t, db.SqliteStoreEngine, conn.Engine()) +} diff --git a/management/internals/shared/db/engine.go b/management/internals/shared/db/engine.go new file mode 100644 index 000000000..02a3f5570 --- /dev/null +++ b/management/internals/shared/db/engine.go @@ -0,0 +1,10 @@ +package db + +// Engine identifies the SQL engine behind a Conn. +type Engine string + +const ( + SqliteStoreEngine Engine = "sqlite" + PostgresStoreEngine Engine = "postgres" + MysqlStoreEngine Engine = "mysql" +) diff --git a/management/internals/shared/db/lock.go b/management/internals/shared/db/lock.go new file mode 100644 index 000000000..bf6fc9b5a --- /dev/null +++ b/management/internals/shared/db/lock.go @@ -0,0 +1,12 @@ +package db + +// LockingStrength is the row lock a query holds until its transaction ends. +type LockingStrength string + +const ( + LockingStrengthUpdate LockingStrength = "UPDATE" + LockingStrengthShare LockingStrength = "SHARE" + LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE" + LockingStrengthKeyShare LockingStrength = "KEY SHARE" + LockingStrengthNone LockingStrength = "NONE" +) diff --git a/management/internals/shared/db/open.go b/management/internals/shared/db/open.go new file mode 100644 index 000000000..48a7bf330 --- /dev/null +++ b/management/internals/shared/db/open.go @@ -0,0 +1,176 @@ +package db + +import ( + "context" + "fmt" + "net/url" + "os" + "path/filepath" + "runtime" + "strings" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "gorm.io/driver/mysql" + "gorm.io/driver/postgres" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +// SqliteFileName is the default SQLite database file inside the data directory. +const SqliteFileName = "store.db" + +// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths +// that bypass gorm. +type PoolConfig struct { + MaxConns int32 + MinConns int32 + MaxConnLifetime time.Duration + HealthCheckPeriod time.Duration +} + +var DefaultPoolConfig = PoolConfig{ + MaxConns: 30, + MinConns: 1, + MaxConnLifetime: 60 * time.Minute, + HealthCheckPeriod: time.Minute, +} + +// GormConfig is the configuration every engine is opened with. +func GormConfig() *gorm.Config { + return &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + CreateBatchSize: 400, + } +} + +// OpenSqlite opens the SQLite database in dataDir, or the file named by +// NB_STORE_ENGINE_SQLITE_FILE. +func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) { + storeFile := SqliteFileName + if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" { + storeFile = envFile + } + return OpenSqliteFile(ctx, dataDir, storeFile) +} + +// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir +// when relative. storeFile may carry SQLite URI query parameters. +func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) { + // Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc") + filePath, query, hasQuery := strings.Cut(storeFile, "?") + + connStr := filePath + if !filepath.IsAbs(filePath) { + connStr = filepath.Join(dataDir, filePath) + } + + // Compose query parameters. User-provided ?_busy_timeout (or its mattn alias + // ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at + // most that long on a lock instead of blocking the only Go-side connection. + // mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so + // the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared + // stays the default on non-Windows for the same reason as before. + parsed, _ := url.ParseQuery(query) + var defaults []string + if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" { + defaults = append(defaults, "_busy_timeout=30000") + } + if !hasQuery && runtime.GOOS != "windows" { + // To avoid `The process cannot access the file because it is being used by another process` on Windows + defaults = append(defaults, "cache=shared") + } + parts := defaults + if hasQuery { + parts = append(parts, query) + } + if len(parts) > 0 { + connStr += "?" + strings.Join(parts, "&") + } + + gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig()) + if err != nil { + return nil, err + } + conn, err := NewConn(ctx, gormDB, SqliteStoreEngine, nil) + if err != nil { + closeGorm(gormDB) + return nil, err + } + return conn, nil +} + +// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool. +func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) { + gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig()) + if err != nil { + return nil, err + } + pgxPool, err := newPgxPool(ctx, dsn, pool) + if err != nil { + closeGorm(gormDB) + return nil, err + } + conn, err := NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool) + if err != nil { + pgxPool.Close() + closeGorm(gormDB) + return nil, err + } + return conn, nil +} + +// MysqlDSN adds the connection parameters every MySQL handle needs, keeping +// the options already present in dsn. +func MysqlDSN(dsn string) string { + separator := "?" + if strings.Contains(dsn, "?") { + separator = "&" + } + return dsn + separator + "charset=utf8&parseTime=True&loc=Local" +} + +// OpenMysql opens a MySQL database through gorm. +func OpenMysql(ctx context.Context, dsn string) (*Conn, error) { + gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig()) + if err != nil { + return nil, err + } + conn, err := NewConn(ctx, gormDB, MysqlStoreEngine, nil) + if err != nil { + closeGorm(gormDB) + return nil, err + } + return conn, nil +} + +func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) { + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + return nil, fmt.Errorf("unable to parse database config: %w", err) + } + + config.MaxConns = cfg.MaxConns + config.MinConns = cfg.MinConns + config.MaxConnLifetime = cfg.MaxConnLifetime + config.HealthCheckPeriod = cfg.HealthCheckPeriod + + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + return nil, fmt.Errorf("unable to create connection pool: %w", err) + } + + if err := pool.Ping(ctx); err != nil { + pool.Close() + return nil, fmt.Errorf("unable to ping database: %w", err) + } + + return pool, nil +} + +func closeGorm(gormDB *gorm.DB) { + if sqlDB, err := gormDB.DB(); err == nil { + _ = sqlDB.Close() + } +} diff --git a/management/internals/shared/db/open_test.go b/management/internals/shared/db/open_test.go new file mode 100644 index 000000000..e6b506b78 --- /dev/null +++ b/management/internals/shared/db/open_test.go @@ -0,0 +1,12 @@ +package db + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestMysqlDSN(t *testing.T) { + assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db")) + assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true")) +} diff --git a/management/internals/shared/db/transaction.go b/management/internals/shared/db/transaction.go new file mode 100644 index 000000000..9699aaee6 --- /dev/null +++ b/management/internals/shared/db/transaction.go @@ -0,0 +1,105 @@ +package db + +import ( + "context" + "errors" + "fmt" + "runtime/debug" + "time" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" +) + +// Tx is an open transaction handed to repository calls; nil means autocommit. +type Tx struct { + db *gorm.DB +} + +// RunInTx runs fn in one transaction that commits when fn returns nil and rolls +// back otherwise, bounded by the configured transaction timeout. +func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error { + timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout) + defer cancel() + + startTime := time.Now() + tx := c.db.WithContext(timeoutCtx).Begin() + if tx.Error != nil { + return tx.Error + } + defer func() { + if r := recover(); r != nil { + tx.Rollback() + panic(r) + } + }() + + if err := c.applyStatementTimeouts(tx); err != nil { + tx.Rollback() + return err + } + + err := c.withForeignKeyChecksDisabled(tx, func() error { + return fn(&Tx{db: tx}) + }) + if err != nil { + tx.Rollback() + c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime) + return err + } + + if err := tx.Commit().Error; err != nil { + c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime) + return err + } + + log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime)) + if c.metrics != nil { + c.metrics.CountTransactionDuration(time.Since(startTime)) + } + return nil +} + +func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error { + if c.engine != PostgresStoreEngine { + return nil + } + if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil { + return fmt.Errorf("failed to set statement timeout: %w", err) + } + if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil { + return fmt.Errorf("failed to set lock timeout: %w", err) + } + return nil +} + +// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids +// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is +// session-scoped and survives a rollback, so it is turned back on whenever fn +// returns or panics; otherwise the pooled connection would keep it disabled. +func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) { + if c.engine != MysqlStoreEngine { + return fn() + } + if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil { + return fmt.Errorf("failed to disable FK checks: %w", err) + } + defer func() { + restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error + if restoreErr == nil { + return + } + if err == nil { + err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr) + return + } + log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr) + }() + return fn() +} + +func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) { + if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) { + log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack()) + } +} diff --git a/management/internals/shared/grpc/loginfilter.go b/management/internals/shared/grpc/loginfilter.go index cc69b7d6e..01ae67cc3 100644 --- a/management/internals/shared/grpc/loginfilter.go +++ b/management/internals/shared/grpc/loginfilter.go @@ -14,6 +14,7 @@ const ( baseBlockDuration = 10 * time.Minute // Duration for which a peer is banned after exceeding the reconnection limit reconnLimitForBan = 30 // Number of reconnections within the reconnTreshold that triggers a ban metaChangeLimit = 5 // Number of reconnections with different metadata that triggers a ban of one peer + maxBanLevel = 6 // Highest ban level; the ban duration doubles per level up to this one ) type lfConfig struct { @@ -21,6 +22,7 @@ type lfConfig struct { baseBlockDuration time.Duration reconnLimitForBan int metaChangeLimit int + maxBanLevel int } func initCfg() *lfConfig { @@ -29,6 +31,7 @@ func initCfg() *lfConfig { baseBlockDuration: baseBlockDuration, reconnLimitForBan: reconnLimitForBan, metaChangeLimit: metaChangeLimit, + maxBanLevel: maxBanLevel, } } @@ -102,11 +105,18 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) { return } - if state.isBanned && now.After(state.banExpiresAt) { + if state.isBanned { + if now.Before(state.banExpiresAt) { + return + } state.isBanned = false } - if state.banLevel > 0 && now.Sub(state.lastSeen) > (2*l.cfg.baseBlockDuration) { + quietSince := state.lastSeen + if state.banExpiresAt.After(quietSince) { + quietSince = state.banExpiresAt + } + if state.banLevel > 0 && now.Sub(quietSince) > (2*l.cfg.baseBlockDuration) { state.banLevel = 0 } @@ -124,10 +134,17 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) { return } + if now.Sub(state.sessionStart) >= l.cfg.reconnThreshold { + state.sessionStart = now + state.sessionCounter = 0 + } + state.sessionCounter++ - if state.sessionCounter > l.cfg.reconnLimitForBan && now.Sub(state.sessionStart) < l.cfg.reconnThreshold { + if state.sessionCounter > l.cfg.reconnLimitForBan { state.isBanned = true - state.banLevel++ + if state.banLevel < l.cfg.maxBanLevel { + state.banLevel++ + } backoffFactor := math.Pow(2, float64(state.banLevel-1)) duration := time.Duration(float64(l.cfg.baseBlockDuration) * backoffFactor) diff --git a/management/internals/shared/grpc/loginfilter_test.go b/management/internals/shared/grpc/loginfilter_test.go index d9df26420..edc256550 100644 --- a/management/internals/shared/grpc/loginfilter_test.go +++ b/management/internals/shared/grpc/loginfilter_test.go @@ -20,6 +20,7 @@ func testAdvancedCfg() *lfConfig { baseBlockDuration: 100 * time.Millisecond, reconnLimitForBan: 3, metaChangeLimit: 2, + maxBanLevel: 3, } } @@ -157,6 +158,187 @@ func (s *LoginFilterTestSuite) TestMetaChangeIsAllowedAfterWindowResets() { s.Equal(1, s.filter.logged[pubKey].metaChangeCounter, "meta change counter should reset") } +func (s *LoginFilterTestSuite) TestReconnectStormAfterQuietPeriodTriggersBan() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + limit := s.filter.cfg.reconnLimitForBan + + s.filter.addLogin(pubKey, meta) + s.Require().Contains(s.filter.logged, pubKey) + s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second)) + + s.filter.addLogin(pubKey, meta) + s.Equal(1, s.filter.logged[pubKey].sessionCounter, "expired window should restart the count") + + for i := 1; i < limit; i++ { + s.filter.addLogin(pubKey, meta) + } + s.True(s.filter.allowLogin(pubKey, meta)) + s.False(s.filter.logged[pubKey].isBanned) + + s.filter.addLogin(pubKey, meta) + + s.False(s.filter.allowLogin(pubKey, meta)) + s.True(s.filter.logged[pubKey].isBanned) +} + +func (s *LoginFilterTestSuite) TestReconnectStormAfterBanExpiresTriggersBanAgain() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + limit := s.filter.cfg.reconnLimitForBan + + for i := 0; i <= limit; i++ { + s.filter.addLogin(pubKey, meta) + } + s.Require().Contains(s.filter.logged, pubKey) + s.Require().True(s.filter.logged[pubKey].isBanned) + + expired := time.Now().Add(-(s.filter.cfg.baseBlockDuration + time.Second)) + s.filter.logged[pubKey].banExpiresAt = expired + s.filter.logged[pubKey].sessionStart = expired + + for i := 0; i <= limit; i++ { + s.filter.addLogin(pubKey, meta) + } + + s.True(s.filter.logged[pubKey].isBanned) + s.Equal(2, s.filter.logged[pubKey].banLevel) +} + +func (s *LoginFilterTestSuite) TestSlowReconnectsAcrossWindowsDoNotBan() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + limit := s.filter.cfg.reconnLimitForBan + + for i := 0; i < limit; i++ { + s.filter.addLogin(pubKey, meta) + } + s.Require().Contains(s.filter.logged, pubKey) + s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second)) + + for i := 0; i < limit; i++ { + s.filter.addLogin(pubKey, meta) + } + + s.True(s.filter.allowLogin(pubKey, meta)) + s.False(s.filter.logged[pubKey].isBanned) +} + +func (s *LoginFilterTestSuite) TestBanLevelEscalatesWhenStormResumesRightAfterBan() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + limit := s.filter.cfg.reconnLimitForBan + banTime := time.Now().Add(-3 * s.filter.cfg.baseBlockDuration) + + s.filter.logged[pubKey] = &peerState{ + currentHash: meta, + isBanned: true, + banLevel: 1, + banExpiresAt: time.Now().Add(-time.Millisecond), + sessionStart: banTime, + lastSeen: banTime, + } + + for i := 0; i <= limit; i++ { + s.filter.addLogin(pubKey, meta) + } + + s.True(s.filter.logged[pubKey].isBanned) + s.Equal(2, s.filter.logged[pubKey].banLevel) +} + +func (s *LoginFilterTestSuite) TestBanLevelResetsAfterQuietPeriodFollowingBan() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + quiet := 2*s.filter.cfg.baseBlockDuration + time.Second + + s.filter.logged[pubKey] = &peerState{ + currentHash: meta, + banLevel: 2, + banExpiresAt: time.Now().Add(-s.filter.cfg.baseBlockDuration), + lastSeen: time.Now().Add(-2 * quiet), + } + + s.filter.addLogin(pubKey, meta) + s.Equal(2, s.filter.logged[pubKey].banLevel, "ban ended more recently than the quiet period") + + s.filter.logged[pubKey].banExpiresAt = time.Now().Add(-quiet) + s.filter.logged[pubKey].lastSeen = time.Now().Add(-2 * quiet) + + s.filter.addLogin(pubKey, meta) + s.Equal(0, s.filter.logged[pubKey].banLevel) +} + +func (s *LoginFilterTestSuite) TestBanDurationIsCappedAtMaxLevel() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + limit := s.filter.cfg.reconnLimitForBan + maxLevel := s.filter.cfg.maxBanLevel + + s.filter.logged[pubKey] = &peerState{ + currentHash: meta, + banLevel: maxLevel, + sessionStart: time.Now(), + lastSeen: time.Now(), + } + + for i := 0; i <= limit; i++ { + s.filter.addLogin(pubKey, meta) + } + + s.True(s.filter.logged[pubKey].isBanned) + s.Equal(maxLevel, s.filter.logged[pubKey].banLevel) + expected := s.filter.cfg.baseBlockDuration << (maxLevel - 1) + s.InDelta(expected, s.filter.logged[pubKey].banExpiresAt.Sub(s.filter.logged[pubKey].lastSeen), float64(time.Millisecond)) +} + +func (s *LoginFilterTestSuite) TestEstablishedPeerReconnectingOnceIsAllowed() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + longAgo := time.Now().Add(-time.Hour) + + s.filter.logged[pubKey] = &peerState{ + currentHash: meta, + sessionCounter: 1, + sessionStart: longAgo, + lastSeen: longAgo, + metaChangeWindowStart: longAgo, + metaChangeCounter: 1, + } + + s.True(s.filter.allowLogin(pubKey, meta)) + s.filter.addLogin(pubKey, meta) + + s.True(s.filter.allowLogin(pubKey, meta)) + s.False(s.filter.logged[pubKey].isBanned) + s.Equal(1, s.filter.logged[pubKey].sessionCounter) +} + +func (s *LoginFilterTestSuite) TestLoginsDuringActiveBanDoNotExtendIt() { + pubKey := "PUB_KEY_A" + meta := uint64(1) + limit := s.filter.cfg.reconnLimitForBan + + for i := 0; i <= limit; i++ { + s.filter.addLogin(pubKey, meta) + } + s.Require().Contains(s.filter.logged, pubKey) + s.Require().True(s.filter.logged[pubKey].isBanned) + expiresAt := time.Now().Add(time.Hour) + s.filter.logged[pubKey].banExpiresAt = expiresAt + lastSeen := s.filter.logged[pubKey].lastSeen + + for i := 0; i <= limit; i++ { + s.filter.addLogin(pubKey, meta) + } + + s.True(s.filter.logged[pubKey].isBanned) + s.Equal(1, s.filter.logged[pubKey].banLevel) + s.Equal(expiresAt, s.filter.logged[pubKey].banExpiresAt) + s.Equal(lastSeen, s.filter.logged[pubKey].lastSeen) + s.Equal(0, s.filter.logged[pubKey].sessionCounter) +} + func BenchmarkHashingMethods(b *testing.B) { meta := nbpeer.PeerSystemMeta{ WtVersion: "1.25.1", diff --git a/management/internals/shared/grpc/pkce_verifier.go b/management/internals/shared/grpc/pkce_verifier.go deleted file mode 100644 index 18155dc1d..000000000 --- a/management/internals/shared/grpc/pkce_verifier.go +++ /dev/null @@ -1,55 +0,0 @@ -package grpc - -import ( - "context" - "fmt" - "time" - - "github.com/eko/gocache/lib/v4/store" - log "github.com/sirupsen/logrus" - - nbcache "github.com/netbirdio/netbird/management/server/cache" -) - -// PKCEVerifierStore manages PKCE verifiers for OAuth flows. -// Supports both in-memory and Redis storage via NB_IDP_CACHE_REDIS_ADDRESS env var. -type PKCEVerifierStore struct { - cache nbcache.Store - ctx context.Context -} - -// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store. -func NewPKCEVerifierStore(ctx context.Context, cacheStore nbcache.Store) *PKCEVerifierStore { - return &PKCEVerifierStore{ - cache: cacheStore, - ctx: ctx, - } -} - -// Store saves a PKCE verifier associated with an OAuth state parameter. -// The verifier is stored with the specified TTL and will be automatically deleted after expiration. -func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) error { - if err := s.cache.Set(s.ctx, state, verifier, store.WithExpiration(ttl)); err != nil { - return fmt.Errorf("failed to store PKCE verifier: %w", err) - } - - log.Debugf("Stored PKCE verifier for state (expires in %s)", ttl) - return nil -} - -// LoadAndDelete retrieves and removes a PKCE verifier for the given state. -// Returns the verifier and true if found, or empty string and false if not found. -// This enforces single-use semantics for PKCE verifiers. -func (s *PKCEVerifierStore) LoadAndDelete(state string) (string, bool) { - verifier, found, err := s.cache.GetDel(s.ctx, state) - if err != nil { - log.Warnf("Failed to consume PKCE verifier: %v", err) - return "", false - } - if !found { - log.Debug("PKCE verifier not found for state") - return "", false - } - - return verifier, true -} diff --git a/management/internals/shared/grpc/proxy.go b/management/internals/shared/grpc/proxy.go index 2fc969ad0..00355087d 100644 --- a/management/internals/shared/grpc/proxy.go +++ b/management/internals/shared/grpc/proxy.go @@ -27,8 +27,6 @@ import ( "google.golang.org/grpc/codes" "google.golang.org/grpc/status" - "github.com/netbirdio/netbird/shared/management/domain" - "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/peers" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" @@ -42,6 +40,7 @@ import ( "github.com/netbirdio/netbird/management/server/users" proxyauth "github.com/netbirdio/netbird/proxy/auth" "github.com/netbirdio/netbird/shared/hash/argon2id" + "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/proto" nbstatus "github.com/netbirdio/netbird/shared/management/status" ) @@ -102,7 +101,8 @@ type ProxyServiceServer struct { mu sync.RWMutex // Manager for reverse proxy operations - serviceManager rpservice.Manager + serviceManager rpservice.Manager + credentialLimits credentialVerificationLimiter // agentNetworkSynth produces synthesised reverse-proxy services from // Agent Network state. Optional — when nil the snapshot path only ships // persisted services. @@ -141,8 +141,8 @@ type ProxyServiceServer struct { // OIDC configuration for proxy authentication oidcConfig ProxyOIDCConfig - // Store for PKCE verifiers - pkceVerifierStore *PKCEVerifierStore + // singleUseStore backs both PKCE verifiers and OIDC session exchange codes. + singleUseStore *SingleUseStore // tokenTTL is the lifetime of one-time tokens generated for proxy // authentication. Defaults to defaultProxyTokenTTL when zero. @@ -157,6 +157,13 @@ type ProxyServiceServer struct { const pkceVerifierTTL = 10 * time.Minute +const sessionCodeTTL = 60 * time.Second + +const sessionCodeCacheNamespace = "proxy:session" + +// The signed nonce binds the handoff mode without changing the state format. +const sessionCodeNoncePrefix = "code." + const defaultProxyTokenTTL = 5 * time.Minute const defaultSnapshotBatchSize = 500 @@ -207,13 +214,13 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error { } // NewProxyServiceServer creates a new proxy service server. -func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer { +func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, singleUseStore *SingleUseStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer { ctx, cancel := context.WithCancel(context.Background()) s := &ProxyServiceServer{ accessLogManager: accessLogMgr, oidcConfig: oidcConfig, tokenStore: tokenStore, - pkceVerifierStore: pkceStore, + singleUseStore: singleUseStore, peersManager: peersManager, usersManager: usersManager, idpManager: idpManager, @@ -242,9 +249,10 @@ func (s *ProxyServiceServer) cleanupStaleProxies(ctx context.Context) { } } -// Close stops background goroutines. +// Close stops background goroutines and releases credential verification state. func (s *ProxyServiceServer) Close() { s.cancel() + s.credentialLimits.close() } // SetServiceManager sets the service manager. Must be called before serving. @@ -304,6 +312,16 @@ func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer { return s.connectAuthorizer } +// GenerateSessionCode creates a single-use code for the given session token. +func (s *ProxyServiceServer) GenerateSessionCode(sessionToken string) (code string, ok bool) { + code, err := s.singleUseStore.Generate(sessionCodeCacheNamespace, sessionToken, sessionCodeTTL) + if err != nil { + log.WithError(err).Error("failed to generate proxy session code") + return "", false + } + return code, true +} + // CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before // forwarding an LLM request upstream. Delegates to the agent-network selector, // which scores applicable policies by remaining headroom and returns the @@ -412,6 +430,7 @@ func (s *ProxyServiceServer) SetProxyController(proxyController proxy.Controller type proxyConnectParams struct { proxyID string address string + version string capabilities *proto.ProxyCapabilities } @@ -422,6 +441,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest return err } params.capabilities = req.GetCapabilities() + params.version = req.GetVersion() conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{ stream: stream, @@ -455,6 +475,7 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings return err } params.capabilities = init.GetCapabilities() + params.version = init.GetVersion() conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{ syncStream: stream, @@ -566,7 +587,7 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params } } - proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, accountID, caps) + proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, params.version, accountID, caps) if err != nil { cancel() if accountID != nil { @@ -1223,6 +1244,7 @@ func shallowCloneMapping(m *proto.ProxyMapping) *proto.ProxyMapping { } } +// Authenticate verifies service credentials and issues a session token. func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) { if err := enforceAccountScope(ctx, req.GetAccountId()); err != nil { return nil, err @@ -1234,6 +1256,14 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen return nil, status.Errorf(codes.FailedPrecondition, "get service from store: %v", err) } + switch req.GetRequest().(type) { + case *proto.AuthenticateRequest_Pin, *proto.AuthenticateRequest_Password: + key := credentialVerificationKey{accountID: credentialAccountID(service.AccountID), serviceID: credentialServiceID(service.ID)} + if err := s.credentialLimits.allow(key); err != nil { + return nil, err + } + } + authenticated, userId, method := s.authenticateRequest(ctx, req, service) // Non-OIDC schemes (PIN/Password/Header) authenticate against per-service @@ -1522,18 +1552,20 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU log.WithContext(ctx).Errorf("failed to get account services: %v", err) return nil, status.Errorf(codes.FailedPrecondition, "get account services: %v", err) } - var found bool + var matchedService *rpservice.Service for _, service := range services { if service.Domain == redirectURL.Hostname() { - found = true + matchedService = service break } } - if !found { + if matchedService == nil { log.WithContext(ctx).Debugf("OIDC redirect URL %q does not match any service domain", redirectURL.Hostname()) return nil, status.Errorf(codes.FailedPrecondition, "service not found in store") } + useSessionCode := s.proxyManager.ClusterSupportsSessionCode(ctx, matchedService.ProxyCluster) + provider, err := oidc.NewProvider(ctx, s.oidcConfig.Issuer) if err != nil { log.WithContext(ctx).Errorf("failed to create OIDC provider: %v", err) @@ -1553,15 +1585,18 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU return nil, status.Errorf(codes.Internal, "generate nonce: %v", err) } nonceB64 := base64.URLEncoding.EncodeToString(nonce) + if useSessionCode { + nonceB64 = sessionCodeNoncePrefix + nonceB64 + } // Using an HMAC here to avoid redirection state being modified. - // State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce) + // State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce) payload := redirectURL.String() + "|" + nonceB64 hmacSum := s.generateHMAC(payload) state := fmt.Sprintf("%s|%s|%s", base64.URLEncoding.EncodeToString([]byte(redirectURL.String())), nonceB64, hmacSum) codeVerifier := oauth2.GenerateVerifier() - if err := s.pkceVerifierStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil { + if err := s.singleUseStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil { log.WithContext(ctx).Errorf("failed to store PKCE verifier: %v", err) return nil, status.Errorf(codes.Internal, "store PKCE verifier: %v", err) } @@ -1598,15 +1633,12 @@ func (s *ProxyServiceServer) generateHMAC(input string) string { return hex.EncodeToString(mac.Sum(nil)) } -// ValidateState validates the state parameter from an OAuth callback. -// Returns the original redirect URL if valid, or an error if invalid. -// The HMAC is verified before consuming the PKCE verifier to prevent -// an attacker from invalidating a legitimate user's auth flow. -func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, err error) { - // State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce) +// ValidateState validates and consumes an OIDC state. +func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, useSessionCode bool, err error) { + // State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce) parts := strings.Split(state, "|") if len(parts) != 3 { - return "", "", errors.New("invalid state format") + return "", "", false, errors.New("invalid state format") } encodedURL := parts[0] @@ -1615,7 +1647,7 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL redirectURLBytes, err := base64.URLEncoding.DecodeString(encodedURL) if err != nil { - return "", "", fmt.Errorf("invalid state encoding: %w", err) + return "", "", false, fmt.Errorf("invalid state encoding: %w", err) } redirectURL = string(redirectURLBytes) @@ -1623,16 +1655,17 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL expectedHMAC := s.generateHMAC(payload) if !hmac.Equal([]byte(providedHMAC), []byte(expectedHMAC)) { - return "", "", errors.New("invalid state signature") + return "", "", false, errors.New("invalid state signature") } + useSessionCode = strings.HasPrefix(nonce, sessionCodeNoncePrefix) // Consume the PKCE verifier only after HMAC validation passes. - verifier, ok := s.pkceVerifierStore.LoadAndDelete(state) + verifier, ok := s.singleUseStore.LoadAndDelete(state) if !ok { - return "", "", errors.New("no verifier for state") + return "", "", false, errors.New("no verifier for state") } - return verifier, redirectURL, nil + return verifier, redirectURL, useSessionCode, nil } // Denied reasons reported to the proxy when access is refused because of the @@ -1837,7 +1870,20 @@ func (s *ProxyServiceServer) getAccountServiceByDomain(ctx context.Context, acco // ValidateSession validates a session token and checks if the user has access to the domain. func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.ValidateSessionRequest) (*proto.ValidateSessionResponse, error) { domain := req.GetDomain() - sessionToken := req.GetSessionToken() + sessionToken := req.GetSessionToken() //nolint:staticcheck + + // A one-time code from the OIDC callback is redeemed here for the durable + // token, so the token never travels in a redirect URL. The redeemed token + // is returned to the proxy (mintedToken) to install as the session cookie. + mintedToken := "" + if code := req.GetSessionCode(); code != "" { + redeemed, found := s.singleUseStore.LoadAndDelete(singleUseCacheKey(sessionCodeCacheNamespace, code)) + if !found { + return deniedSessionResponse("invalid or expired session code"), nil + } + sessionToken = redeemed + mintedToken = redeemed + } if domain == "" || sessionToken == "" { return deniedSessionResponse("missing domain or session_token"), nil @@ -1910,6 +1956,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val UserEmail: user.Email, PeerGroupIds: groupIDs, PeerGroupNames: groupNames, + SessionToken: mintedToken, }, nil } diff --git a/management/internals/shared/grpc/proxy_connect_version_test.go b/management/internals/shared/grpc/proxy_connect_version_test.go new file mode 100644 index 000000000..e2fc49a38 --- /dev/null +++ b/management/internals/shared/grpc/proxy_connect_version_test.go @@ -0,0 +1,93 @@ +package grpc + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/shared/management/proto" +) + +const ( + versionTestProxyID = "proxy-a" + versionTestCluster = "cluster.example.com" + versionTestVersion = "0.60.0" +) + +// hangupStream cancels its context on the first Send, emulating a proxy that +// disconnects right after receiving the initial snapshot. The legacy stream +// carries no proxy-to-management messages, so this is the only way for +// GetMappingUpdate to return. +type hangupStream struct { + recordingStream + ctx context.Context + cancel context.CancelFunc +} + +func (s *hangupStream) Send(m *proto.GetMappingUpdateResponse) error { + s.cancel() + return s.recordingStream.Send(m) +} + +func (s *hangupStream) Context() context.Context { return s.ctx } + +// newVersionTestServer wires a server whose proxy manager only accepts a +// Connect carrying versionTestVersion, so a dropped or mangled version fails +// the test as an unexpected call. +func newVersionTestServer(t *testing.T) *ProxyServiceServer { + t.Helper() + ctrl := gomock.NewController(t) + + svcMgr := rpservice.NewMockManager(ctrl) + svcMgr.EXPECT().GetGlobalServices(gomock.Any()).Return(nil, nil) + + proxyMgr := proxy.NewMockManager(ctrl) + proxyMgr.EXPECT(). + Connect(gomock.Any(), versionTestProxyID, gomock.Any(), versionTestCluster, gomock.Any(), versionTestVersion, gomock.Any(), gomock.Any()). + Return(&proxy.Proxy{ID: versionTestProxyID, Version: versionTestVersion}, nil) + proxyMgr.EXPECT().Disconnect(gomock.Any(), versionTestProxyID, gomock.Any()).Return(nil) + + s := newSnapshotTestServer(t, 10) + s.serviceManager = svcMgr + s.proxyManager = proxyMgr + return s +} + +func TestSyncMappings_ForwardsProxyVersion(t *testing.T) { + s := newVersionTestServer(t) + + // The init carries the version, the ack acknowledges the empty snapshot, + // and the exhausted fake stream then ends the RPC. + stream := &syncRecordingStream{ + recvMsgs: []*proto.SyncMappingsRequest{ + {Msg: &proto.SyncMappingsRequest_Init{Init: &proto.SyncMappingsInit{ + ProxyId: versionTestProxyID, + Address: versionTestCluster, + Version: versionTestVersion, + }}}, + ackMsg(), + }, + } + + err := s.SyncMappings(stream) + require.ErrorContains(t, err, "no more recv messages") +} + +func TestGetMappingUpdate_ForwardsProxyVersion(t *testing.T) { + s := newVersionTestServer(t) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + stream := &hangupStream{ctx: ctx, cancel: cancel} + + err := s.GetMappingUpdate(&proto.GetMappingUpdateRequest{ + ProxyId: versionTestProxyID, + Address: versionTestCluster, + Version: versionTestVersion, + }, stream) + require.ErrorIs(t, err, context.Canceled) +} diff --git a/management/internals/shared/grpc/proxy_credential_limiter.go b/management/internals/shared/grpc/proxy_credential_limiter.go new file mode 100644 index 000000000..2a3a02347 --- /dev/null +++ b/management/internals/shared/grpc/proxy_credential_limiter.go @@ -0,0 +1,101 @@ +package grpc + +import ( + "sync" + "time" + + "golang.org/x/time/rate" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "google.golang.org/protobuf/types/known/durationpb" +) + +const ( + credentialVerificationInterval = 6 * time.Second + credentialVerificationBurst = 5 + credentialVerificationMaxServices = 4096 + credentialVerificationIdleTimeout = 15 * time.Minute + credentialVerificationCleanupInterval = time.Minute +) + +type credentialAccountID string +type credentialServiceID string + +type credentialVerificationKey struct { + accountID credentialAccountID + serviceID credentialServiceID +} + +type credentialVerificationBudget struct { + limiter *rate.Limiter + lastUsed time.Time +} + +// The zero value is ready to use. Budgets are local to this Management process; +// proxy replicas reaching this process share a service's verification budget. +type credentialVerificationLimiter struct { + mu sync.Mutex + now func() time.Time + services map[credentialVerificationKey]*credentialVerificationBudget + nextCleanup time.Time + closed bool +} + +func (l *credentialVerificationLimiter) allow(key credentialVerificationKey) error { + l.mu.Lock() + defer l.mu.Unlock() + if l.closed { + return status.Error(codes.Unavailable, "credential verification is closed") + } + now := time.Now() + if l.now != nil { + now = l.now() + } + l.cleanup(now) + budget := l.services[key] + if budget == nil { + if len(l.services) >= credentialVerificationMaxServices { + return credentialVerificationThrottled(credentialVerificationCleanupInterval) + } + if l.services == nil { + l.services = make(map[credentialVerificationKey]*credentialVerificationBudget) + } + budget = &credentialVerificationBudget{limiter: rate.NewLimiter(rate.Every(credentialVerificationInterval), credentialVerificationBurst)} + l.services[key] = budget + } + budget.lastUsed = now + if budget.limiter.AllowN(now, 1) { + return nil + } + delay := max(time.Nanosecond, time.Duration((1-budget.limiter.TokensAt(now))*float64(credentialVerificationInterval))) + return credentialVerificationThrottled(delay) +} + +func (l *credentialVerificationLimiter) cleanup(now time.Time) { + if now.Before(l.nextCleanup) { + return + } + l.nextCleanup = now.Add(credentialVerificationCleanupInterval) + for key, budget := range l.services { + if now.Sub(budget.lastUsed) >= credentialVerificationIdleTimeout { + delete(l.services, key) + } + } +} + +func (l *credentialVerificationLimiter) close() { + l.mu.Lock() + defer l.mu.Unlock() + l.closed = true + l.services = nil +} + +func credentialVerificationThrottled(delay time.Duration) error { + s := status.New(codes.ResourceExhausted, "too many credential verification attempts") + withRetry, err := s.WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(delay)}) + if err != nil { + return s.Err() + } + return withRetry.Err() +} diff --git a/management/internals/shared/grpc/proxy_credential_limiter_test.go b/management/internals/shared/grpc/proxy_credential_limiter_test.go new file mode 100644 index 000000000..565b14889 --- /dev/null +++ b/management/internals/shared/grpc/proxy_credential_limiter_test.go @@ -0,0 +1,79 @@ +package grpc + +import ( + "strconv" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func TestCredentialVerificationRefillAndIsolation(t *testing.T) { + now := time.Now() + l := credentialVerificationLimiter{now: func() time.Time { return now }} + key := credentialVerificationKey{accountID: "account", serviceID: "service"} + for range credentialVerificationBurst { + require.NoError(t, l.allow(key)) + } + err := l.allow(key) + require.Equal(t, codes.ResourceExhausted, status.Code(err), "the burst must be bounded") + now = now.Add(3 * time.Second) + err = l.allow(key) + require.Equal(t, codes.ResourceExhausted, status.Code(err), "a partially refilled token must not permit a check") + details := status.Convert(err).Details() + require.Len(t, details, 1, "throttling must provide RetryInfo") + retry, ok := details[0].(*errdetails.RetryInfo) + require.True(t, ok, "retry details must use the standard message") + assert.Equal(t, 3*time.Second, retry.RetryDelay.AsDuration(), "retry hint must reflect time until the next check") + now = now.Add(3 * time.Second) + require.NoError(t, l.allow(key)) + assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "only one check must refill every six seconds") + require.NoError(t, l.allow(credentialVerificationKey{accountID: "other-account", serviceID: key.serviceID})) + require.NoError(t, l.allow(credentialVerificationKey{accountID: key.accountID, serviceID: "other-service"})) +} + +func TestCredentialVerificationCapacityAndExpiry(t *testing.T) { + now := time.Now() + l := credentialVerificationLimiter{now: func() time.Time { return now }} + for i := range credentialVerificationMaxServices { + require.NoError(t, l.allow(credentialVerificationKey{accountID: "account", serviceID: credentialServiceID(strconv.Itoa(i))})) + } + key := credentialVerificationKey{accountID: "account", serviceID: "new-service"} + assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "capacity exhaustion must deny new checks") + now = now.Add(credentialVerificationIdleTimeout) + for range credentialVerificationBurst { + require.NoError(t, l.allow(key)) + } + assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "expiry must retain the normal burst bound") +} + +func TestCredentialVerificationConcurrentChecksAndClose(t *testing.T) { + var l credentialVerificationLimiter + key := credentialVerificationKey{accountID: "account", serviceID: "service"} + var admitted atomic.Int32 + var wg sync.WaitGroup + for range 100 { + wg.Go(func() { + if err := l.allow(key); err == nil { + admitted.Add(1) + } else { + assert.Equal(t, codes.ResourceExhausted, status.Code(err), "excess checks must be throttled") + } + }) + } + wg.Wait() + assert.EqualValues(t, credentialVerificationBurst, admitted.Load(), "concurrent checks must share the burst") + for range 10 { + wg.Go(l.close) + wg.Go(func() { assert.Error(t, l.allow(key)) }) + } + wg.Wait() + assert.Empty(t, l.services, "closing must release retained budgets") + assert.Equal(t, codes.Unavailable, status.Code(l.allow(key)), "checks after close must fail closed") +} diff --git a/management/internals/shared/grpc/proxy_credentials.md b/management/internals/shared/grpc/proxy_credentials.md new file mode 100644 index 000000000..1caf432e5 --- /dev/null +++ b/management/internals/shared/grpc/proxy_credentials.md @@ -0,0 +1,18 @@ +# Reverse proxy credential verification + +The `ProxyService.Authenticate` RPC limits PIN and password checks before +verifying their Argon2 hashes. Both methods share one budget per account and +service: a burst of five checks, replenishing one check every six seconds +(ten per minute). Successful and failed checks consume the budget. Account +scope and service lookup run before the limiter. + +Excess checks receive gRPC `ResourceExhausted` with a standard `RetryInfo` delay. +Updated proxies translate it to HTTP 429 and `Retry-After`. Older proxies show +an authentication-service error but cannot bypass the Management limit. + +Budgets are held in memory per Management process and reset on restart. Proxy +replicas reaching the same Management process share its budgets. Multiple +Management processes have independent budgets; this is not a cluster-wide +limit. At most 4,096 service budgets are retained, with idle entries expiring +after fifteen minutes. Capacity exhaustion denies new checks until entries +expire. Closing the server releases the retained state. diff --git a/management/internals/shared/grpc/proxy_credentials_test.go b/management/internals/shared/grpc/proxy_credentials_test.go new file mode 100644 index 000000000..2e144bd69 --- /dev/null +++ b/management/internals/shared/grpc/proxy_credentials_test.go @@ -0,0 +1,131 @@ +package grpc_test + +import ( + "context" + "net" + "net/netip" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/peer" + "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey" + nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/proto" +) + +func credentialServer(t *testing.T) (*nbgrpc.ProxyServiceServer, context.Context, grpc.UnaryServerInterceptor) { + t.Helper() + ctx := context.Background() + s, err := store.NewStore(ctx, types.SqliteStoreEngine, t.TempDir(), nil, false) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) }) + require.NoError(t, s.SaveAccount(ctx, &types.Account{Id: "account"})) + keys, err := sessionkey.GenerateKeyPair() + require.NoError(t, err) + for _, id := range []string{"service", "other-service"} { + svc := &service.Service{ + ID: id, AccountID: "account", Name: id, Domain: id + ".example.com", + Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey, + Auth: service.AuthConfig{ + PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"}, + PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "test-password"}, + }, + } + require.NoError(t, svc.Auth.HashSecrets()) + require.NoError(t, s.CreateService(ctx, svc)) + } + account := "account" + token, err := types.CreateNewProxyAccessToken("test proxy", time.Hour, &account, "admin") + require.NoError(t, err) + require.NoError(t, s.SaveProxyAccessToken(ctx, &token.ProxyAccessToken)) + ctx = metadata.NewIncomingContext(ctx, metadata.Pairs("authorization", "Bearer "+string(token.PlainToken))) + ctx = peer.NewContext(ctx, &peer.Peer{Addr: net.TCPAddrFromAddrPort(netip.MustParseAddrPort("192.0.2.1:443"))}) + server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) + t.Cleanup(server.Close) + server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil)) + interceptor, _, closeInterceptor := nbgrpc.NewProxyAuthInterceptors(s) + t.Cleanup(closeInterceptor) + return server, ctx, interceptor +} + +func TestAuthenticateCredentialRateLimit(t *testing.T) { + server, ctx, interceptor := credentialServer(t) + authenticate := func(req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) { + response, err := interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: "/management.ProxyService/Authenticate"}, func(ctx context.Context, req any) (any, error) { + return server.Authenticate(ctx, req.(*proto.AuthenticateRequest)) + }) + if err != nil { + return nil, err + } + return response.(*proto.AuthenticateResponse), nil + } + for i := range 5 { + req := &proto.AuthenticateRequest{AccountId: "account", Id: "service"} + if i%2 == 0 { + req.Request = &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}} + } else { + req.Request = &proto.AuthenticateRequest_Password{Password: &proto.PasswordRequest{Password: "wrong-password"}} + } + resp, err := authenticate(req) + require.NoError(t, err) + assert.False(t, resp.GetSuccess(), "incorrect PINs and passwords must be denied") + assert.Empty(t, resp.GetSessionToken(), "incorrect credentials must not issue a token") + } + req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "842716"}}} + resp, err := authenticate(req) + assert.Nil(t, resp, "a throttled verification must not return a session") + require.Equal(t, codes.ResourceExhausted, status.Code(err), "PIN and password checks must share a service budget even with a valid proxy token") + details := status.Convert(err).Details() + require.Len(t, details, 1, "throttled responses must include a retry hint") + retry, ok := details[0].(*errdetails.RetryInfo) + require.True(t, ok, "the hint must use the standard RetryInfo message") + assert.Positive(t, retry.RetryDelay.AsDuration(), "the retry delay must be positive") + assert.LessOrEqual(t, retry.RetryDelay.AsDuration(), 6*time.Second, "the service must replenish one verification every six seconds") + req.AccountId = "another-account" + _, err = authenticate(req) + assert.Equal(t, codes.PermissionDenied, status.Code(err), "account scope must still be enforced before throttling") + req.AccountId = "account" + req.Id = "other-service" + resp, err = authenticate(req) + require.NoError(t, err) + assert.True(t, resp.GetSuccess(), "one service's throttle must not block another service") + assert.NotEmpty(t, resp.GetSessionToken(), "valid credentials on another service must issue a session") +} + +func TestAuthenticateCredentialConcurrentLimit(t *testing.T) { + server, _, _ := credentialServer(t) + req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}} + var checked, throttled atomic.Int32 + var wg sync.WaitGroup + for range 20 { + wg.Go(func() { + resp, err := server.Authenticate(context.Background(), req) + switch status.Code(err) { + case codes.OK: + checked.Add(1) + assert.False(t, resp.GetSuccess(), "incorrect credentials must be denied") + case codes.ResourceExhausted: + throttled.Add(1) + default: + assert.NoError(t, err) + } + }) + } + wg.Wait() + assert.EqualValues(t, 5, checked.Load(), "only the burst budget may reach concurrent credential verification") + assert.EqualValues(t, 15, throttled.Load(), "excess concurrent checks must be throttled") +} diff --git a/management/internals/shared/grpc/proxy_test.go b/management/internals/shared/grpc/proxy_test.go index 29b7c9523..1e718a2c4 100644 --- a/management/internals/shared/grpc/proxy_test.go +++ b/management/internals/shared/grpc/proxy_test.go @@ -129,11 +129,11 @@ func drainEmpty(ch chan *proto.GetMappingUpdateResponse) bool { func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) { ctx := context.Background() tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t)) - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) s := &ProxyServiceServer{ - tokenStore: tokenStore, - pkceVerifierStore: pkceStore, + tokenStore: tokenStore, + singleUseStore: singleUseStore, } s.SetProxyController(newTestProxyController()) @@ -186,11 +186,11 @@ func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) { func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) { ctx := context.Background() tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t)) - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) s := &ProxyServiceServer{ - tokenStore: tokenStore, - pkceVerifierStore: pkceStore, + tokenStore: tokenStore, + singleUseStore: singleUseStore, } s.SetProxyController(newTestProxyController()) @@ -220,11 +220,11 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) { func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) { ctx := context.Background() tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t)) - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) s := &ProxyServiceServer{ - tokenStore: tokenStore, - pkceVerifierStore: pkceStore, + tokenStore: tokenStore, + singleUseStore: singleUseStore, } s.SetProxyController(newTestProxyController()) @@ -272,13 +272,13 @@ func generateState(s *ProxyServiceServer, redirectURL string) string { func TestOAuthState_NeverTheSame(t *testing.T) { ctx := context.Background() - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) s := &ProxyServiceServer{ oidcConfig: ProxyOIDCConfig{ HMACKey: []byte("test-hmac-key"), }, - pkceVerifierStore: pkceStore, + singleUseStore: singleUseStore, } redirectURL := "https://app.example.com/callback" @@ -300,20 +300,20 @@ func TestOAuthState_NeverTheSame(t *testing.T) { func TestValidateState_RejectsOldTwoPartFormat(t *testing.T) { ctx := context.Background() - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) s := &ProxyServiceServer{ oidcConfig: ProxyOIDCConfig{ HMACKey: []byte("test-hmac-key"), }, - pkceVerifierStore: pkceStore, + singleUseStore: singleUseStore, } // Old format had only 2 parts: base64(url)|hmac - err := s.pkceVerifierStore.Store("base64url|hmac", "test", 10*time.Minute) + err := s.singleUseStore.Store("base64url|hmac", "test", 10*time.Minute) require.NoError(t, err) - _, _, err = s.ValidateState("base64url|hmac") + _, _, _, err = s.ValidateState("base64url|hmac") require.Error(t, err) assert.Contains(t, err.Error(), "invalid state format") } @@ -372,24 +372,48 @@ func TestEnforceAccountScope_AllowsNoTokenInContext(t *testing.T) { func TestValidateState_RejectsInvalidHMAC(t *testing.T) { ctx := context.Background() - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) s := &ProxyServiceServer{ oidcConfig: ProxyOIDCConfig{ HMACKey: []byte("test-hmac-key"), }, - pkceVerifierStore: pkceStore, + singleUseStore: singleUseStore, } // Store with tampered HMAC - err := s.pkceVerifierStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute) + err := s.singleUseStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute) require.NoError(t, err) - _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac") + _, _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac") require.Error(t, err) assert.Contains(t, err.Error(), "invalid state signature") } +func TestSessionCodeCannotConsumeOIDCState(t *testing.T) { + const verifier = "pkce-verifier" + + store := NewSingleUseStore(context.Background(), testCacheStore(t)) + server := &ProxyServiceServer{ + oidcConfig: ProxyOIDCConfig{ + HMACKey: []byte("test-hmac-key"), + }, + singleUseStore: store, + } + state := generateState(server, "https://service.example.com/callback") + require.NoError(t, store.Store(state, verifier, time.Minute)) + + response, err := server.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ + SessionCode: state, + }) + require.NoError(t, err) + assert.False(t, response.GetValid()) + + gotVerifier, _, _, err := server.ValidateState(state) + require.NoError(t, err) + assert.Equal(t, verifier, gotVerifier) +} + func TestSendServiceUpdateToCluster_FiltersOnCapability(t *testing.T) { tokenStore := NewOneTimeTokenStore(context.Background(), testCacheStore(t)) diff --git a/management/internals/shared/grpc/single_use.go b/management/internals/shared/grpc/single_use.go new file mode 100644 index 000000000..467ee6020 --- /dev/null +++ b/management/internals/shared/grpc/single_use.go @@ -0,0 +1,67 @@ +package grpc + +import ( + "context" + "crypto/rand" + "encoding/base64" + "fmt" + "time" + + "github.com/eko/gocache/lib/v4/store" + log "github.com/sirupsen/logrus" + + nbcache "github.com/netbirdio/netbird/management/server/cache" +) + +// SingleUseStore stores short-lived values that can be retrieved only once. +type SingleUseStore struct { + cache nbcache.Store + ctx context.Context +} + +// NewSingleUseStore creates a single-use value store over the shared cache. +func NewSingleUseStore(ctx context.Context, cacheStore nbcache.Store) *SingleUseStore { + return &SingleUseStore{ + cache: cacheStore, + ctx: ctx, + } +} + +// Store saves value under key with the given TTL, after which it is evicted. +func (s *SingleUseStore) Store(key, value string, ttl time.Duration) error { + if err := s.cache.Set(s.ctx, key, value, store.WithExpiration(ttl)); err != nil { + return fmt.Errorf("store single-use value: %w", err) + } + return nil +} + +// Generate stores a value under a namespaced random key and returns the random key. +func (s *SingleUseStore) Generate(namespace, value string, ttl time.Duration) (string, error) { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return "", fmt.Errorf("generate single-use key: %w", err) + } + + key := base64.RawURLEncoding.EncodeToString(buf) + if err := s.Store(singleUseCacheKey(namespace, key), value, ttl); err != nil { + return "", err + } + return key, nil +} + +func singleUseCacheKey(namespace, key string) string { + return namespace + ":" + key +} + +// LoadAndDelete retrieves and removes the value for a key. +func (s *SingleUseStore) LoadAndDelete(key string) (string, bool) { + value, found, err := s.cache.GetDel(s.ctx, key) + if err != nil { + log.Warnf("failed to consume single-use value: %v", err) + return "", false + } + if !found { + return "", false + } + return value, true +} diff --git a/management/internals/shared/grpc/pkce_verifier_test.go b/management/internals/shared/grpc/single_use_test.go similarity index 59% rename from management/internals/shared/grpc/pkce_verifier_test.go rename to management/internals/shared/grpc/single_use_test.go index e7175b6c5..8243ca43b 100644 --- a/management/internals/shared/grpc/pkce_verifier_test.go +++ b/management/internals/shared/grpc/single_use_test.go @@ -6,7 +6,7 @@ import ( "time" ) -func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) { +func TestSingleUseStoreLoadAndDelete(t *testing.T) { const ( state = "state" verifier = "verifier" @@ -14,7 +14,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) { ) t.Run("exactly one concurrent caller consumes the verifier", func(t *testing.T) { - store := NewPKCEVerifierStore(context.Background(), testCacheStore(t)) + store := NewSingleUseStore(context.Background(), testCacheStore(t)) if err := store.Store(state, verifier, time.Minute); err != nil { t.Fatalf("couldn't store PKCE verifier: %s", err) } @@ -50,7 +50,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) { }) t.Run("replayed state is rejected", func(t *testing.T) { - store := NewPKCEVerifierStore(context.Background(), testCacheStore(t)) + store := NewSingleUseStore(context.Background(), testCacheStore(t)) if err := store.Store(state, verifier, time.Minute); err != nil { t.Fatalf("couldn't store PKCE verifier: %s", err) } @@ -64,7 +64,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) { }) t.Run("unknown state is rejected", func(t *testing.T) { - store := NewPKCEVerifierStore(context.Background(), testCacheStore(t)) + store := NewSingleUseStore(context.Background(), testCacheStore(t)) if got, found := store.LoadAndDelete("never-stored"); found { t.Fatalf("unknown state should not resolve, got %q", got) @@ -72,7 +72,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) { }) t.Run("expired verifier is rejected", func(t *testing.T) { - store := NewPKCEVerifierStore(context.Background(), testCacheStore(t)) + store := NewSingleUseStore(context.Background(), testCacheStore(t)) if err := store.Store(state, verifier, 50*time.Millisecond); err != nil { t.Fatalf("couldn't store PKCE verifier: %s", err) } @@ -83,3 +83,40 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) { } }) } + +func TestSingleUseStore_GenerateAndConsumeOnce(t *testing.T) { + const namespace = "test" + s := NewSingleUseStore(context.Background(), testCacheStore(t)) + + key, err := s.Generate(namespace, "the-value", time.Minute) + if err != nil { + t.Fatalf("generate: %v", err) + } + if key == "" || key == "the-value" { + t.Fatalf("unexpected key %q", key) + } + + value, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)) + if !found || value != "the-value" { + t.Fatalf("expected to load the stored value, got %q found=%v", value, found) + } + + if _, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)); found { + t.Fatal("value must be consumed on first LoadAndDelete") + } +} + +func TestSingleUseStore_GenerateUniqueKeys(t *testing.T) { + s := NewSingleUseStore(context.Background(), testCacheStore(t)) + a, err := s.Generate("test", "v", time.Minute) + if err != nil { + t.Fatalf("generate a: %v", err) + } + b, err := s.Generate("test", "v", time.Minute) + if err != nil { + t.Fatalf("generate b: %v", err) + } + if a == b { + t.Fatal("generated keys must be distinct") + } +} diff --git a/management/internals/shared/grpc/validate_session_test.go b/management/internals/shared/grpc/validate_session_test.go index 4e70e61e4..4b36e74cf 100644 --- a/management/internals/shared/grpc/validate_session_test.go +++ b/management/internals/shared/grpc/validate_session_test.go @@ -40,9 +40,9 @@ func setupValidateSessionTest(t *testing.T) *validateSessionTestSetup { proxyManager := &testValidateSessionProxyManager{} tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t)) - pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t)) + singleUseStore := NewSingleUseStore(ctx, testCacheStore(t)) - proxyService := NewProxyServiceServer(nil, tokenStore, pkceStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil) + proxyService := NewProxyServiceServer(nil, tokenStore, singleUseStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil) proxyService.SetServiceManager(serviceManager) createTestProxies(t, ctx, testStore) @@ -570,7 +570,7 @@ func (m *testValidateSessionServiceManager) DeleteAccountCluster(_ context.Conte type testValidateSessionProxyManager struct{} -func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) { +func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) { return nil, nil } @@ -634,6 +634,10 @@ func (m *testValidateSessionProxyManager) ClusterSupportsPrivate(_ context.Conte return nil } +func (m *testValidateSessionProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool { + return false +} + type testValidateSessionUsersManager struct { store store.Store } @@ -662,3 +666,47 @@ func (m *testValidateSessionUsersManager) GetUserWithGroups(ctx context.Context, } return user, groups, nil } + +func TestValidateSession_RedeemsSessionCode(t *testing.T) { + setup := setupValidateSessionTest(t) + defer setup.cleanup() + + proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "testProxyId") + require.NoError(t, err) + + token := createSessionToken(t, proxy.SessionPrivateKey, "allowedUserId", "test-proxy.example.com") + code, ok := setup.proxyService.GenerateSessionCode(token) + require.True(t, ok) + require.NotEqual(t, token, code, "code must not be the token itself") + + resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ + Domain: "test-proxy.example.com", + SessionCode: code, + }) + require.NoError(t, err) + assert.True(t, resp.Valid, "redeemed code should authorize the user") + assert.Equal(t, "allowedUserId", resp.UserId) + assert.Equal(t, token, resp.GetSessionToken(), "response must carry the durable token for the cookie") + + // Single-use: the same code must not redeem again. + resp2, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ + Domain: "test-proxy.example.com", + SessionCode: code, + }) + require.NoError(t, err) + assert.False(t, resp2.Valid, "a consumed code must be rejected") + assert.Empty(t, resp2.GetSessionToken()) +} + +func TestValidateSession_InvalidSessionCode(t *testing.T) { + setup := setupValidateSessionTest(t) + defer setup.cleanup() + + resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ + Domain: "test-proxy.example.com", + SessionCode: "does-not-exist", + }) + require.NoError(t, err) + assert.False(t, resp.Valid) + assert.Empty(t, resp.GetSessionToken()) +} diff --git a/management/server/account.go b/management/server/account.go index 6ccf673f5..038c5d8db 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -112,6 +112,9 @@ type DefaultAccountManager struct { permissionsManager permissions.Manager disableDefaultPolicy bool + + deletionHooksMu sync.RWMutex + deletionHooks []account.DeletionHook } var _ account.Manager = (*DefaultAccountManager)(nil) @@ -120,6 +123,32 @@ func (am *DefaultAccountManager) SetServiceManager(serviceManager service.Manage am.serviceManager = serviceManager } +// AddAccountDeletionHook registers hook to run on every account deletion. Hooks run in +// registration order, and the first one to fail stops the rest and aborts the deletion. +// It panics on a nil hook: dropping one silently would skip that hook's cleanup on every +// deletion, so the wiring bug surfaces at startup instead. +func (am *DefaultAccountManager) AddAccountDeletionHook(hook account.DeletionHook) { + if hook == nil { + panic("nil account deletion hook") + } + am.deletionHooksMu.Lock() + defer am.deletionHooksMu.Unlock() + am.deletionHooks = append(am.deletionHooks, hook) +} + +func (am *DefaultAccountManager) runAccountDeletionHooks(ctx context.Context, accountID string) error { + am.deletionHooksMu.RLock() + hooks := slices.Clone(am.deletionHooks) + am.deletionHooksMu.RUnlock() + + for _, hook := range hooks { + if err := hook(ctx, accountID); err != nil { + return fmt.Errorf("account deletion hook: %w", err) + } + } + return nil +} + func isUniqueConstraintError(err error) bool { switch { case strings.Contains(err.Error(), "(SQLSTATE 23505)"), @@ -305,6 +334,9 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco var groupChangesAffectPeers bool var reloadReverseProxy bool var effectiveOldNetworkRange netip.Prefix + var ipv6Changed bool + var ipv6Snap *affectedpeers.Snapshot + var ipv6Change affectedpeers.Change err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { var groupsUpdated bool @@ -350,10 +382,10 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco } if ipv6SettingsChanged(oldSettings, newSettings) { - if err = am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings); err != nil { + if ipv6Change, err = am.applyIPv6SettingsChange(ctx, transaction, accountID, oldSettings, newSettings); err != nil { return err } - updateAccountPeers = true + ipv6Changed = true } if oldSettings.RoutingPeerDNSResolutionEnabled != newSettings.RoutingPeerDNSResolutionEnabled || @@ -390,12 +422,20 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco return err } - if updateAccountPeers || groupsUpdated { + if updateAccountPeers || groupsUpdated || ipv6Changed { if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { return err } } + // A full account refresh already covers the IPv6 change, so the affected-peers + // snapshot is only needed when nothing account-wide changed. + if ipv6Changed && !updateAccountPeers && !groupChangesAffectPeers { + if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + } + return nil }) if err != nil { @@ -457,13 +497,34 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco } } - if updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers { - go am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate}) + switch { + case updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers: + go am.UpdateAccountPeers(context.WithoutCancel(ctx), accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate}) + case ipv6Snap != nil: + am.ExpandAndUpdateAffected(ctx, accountID, ipv6Snap, ipv6Change) } return newSettings, nil } +// applyIPv6SettingsChange reconciles peer IPv6 addresses for new IPv6 settings and +// returns the affected-peers change: peers whose address changed refresh together +// with every peer that reaches them. On a range change every peer holding an address +// also refreshes itself, since its interface prefix comes from the account range even +// when its address stays inside the new one. +func (am *DefaultAccountManager) applyIPv6SettingsChange(ctx context.Context, transaction store.Store, accountID string, oldSettings, newSettings *types.Settings) (affectedpeers.Change, error) { + result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings) + if err != nil { + return affectedpeers.Change{}, err + } + + change := affectedpeers.Change{ChangedPeerIDs: result.changed} + if oldSettings.NetworkRangeV6 != newSettings.NetworkRangeV6 { + change.OutputPeerIDs = result.withIPv6 + } + return change, nil +} + func ipv6SettingsChanged(old, updated *types.Settings) bool { if old.NetworkRangeV6 != updated.NetworkRangeV6 { return true @@ -889,6 +950,10 @@ func (am *DefaultAccountManager) DeleteAccount(ctx context.Context, accountID, u return status.Errorf(status.Internal, "failed to build user infos for account %s: %v", accountID, err) } + if err = am.runAccountDeletionHooks(ctx, accountID); err != nil { + return err + } + if err = am.deleteAccountUsers(ctx, accountID, userID, account.Users, userInfosMap); err != nil { return err } @@ -1709,9 +1774,11 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth change.LinkGroups = allGroupChanges - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges) + if err != nil { return fmt.Errorf("reconcile IPv6 for group changes: %w", err) } + change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...) if err = transaction.IncrementNetworkSerial(ctx, userAuth.AccountId); err != nil { return fmt.Errorf("error incrementing network serial: %w", err) @@ -2301,7 +2368,8 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte return false, false, err } - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups) + if err != nil { return false, false, fmt.Errorf("reconcile IPv6 for group changes: %w", err) } @@ -2310,7 +2378,7 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte return false, false, fmt.Errorf("error checking if group changes affect peers: %w", err) } - return len(updatedGroups) > 0, peersAffected, nil + return len(updatedGroups) > 0, peersAffected || len(ipv6Changed) > 0, nil } // propagateAutoGroupsForUsers adds each user's peers to their AutoGroups where not already present. @@ -2358,7 +2426,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t return err } - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "") if err != nil { return err } @@ -2395,7 +2463,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t // v6 address get one allocated. When disabled, all v6 addresses are cleared. // When the v6 range changes, all v6 addresses are reallocated. func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transaction store.Store, accountID, peerID string, newIPv6 netip.Addr) error { - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "") if err != nil { return fmt.Errorf("get peers: %w", err) } @@ -2407,56 +2475,78 @@ func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transac return nil } -func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) error { - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "") +// ipv6Reassignment reports the outcome of an IPv6 address reconciliation. +type ipv6Reassignment struct { + // changed are the peers whose IPv6 address was assigned, removed or reallocated. + changed []string + // withIPv6 are all peers holding an IPv6 address after the reconciliation. + withIPv6 []string +} + +func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) (ipv6Reassignment, error) { + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "") if err != nil { - return fmt.Errorf("get peers: %w", err) + return ipv6Reassignment{}, fmt.Errorf("get peers: %w", err) } network, err := transaction.GetAccountNetwork(ctx, store.LockingStrengthUpdate, accountID) if err != nil { - return fmt.Errorf("get network: %w", err) + return ipv6Reassignment{}, fmt.Errorf("get network: %w", err) } if err := am.ensureIPv6Subnet(ctx, transaction, accountID, settings, network); err != nil { - return err + return ipv6Reassignment{}, err } allowedPeers, err := am.buildIPv6AllowedPeers(ctx, transaction, accountID, settings) if err != nil { - return err + return ipv6Reassignment{}, err } v6Prefix, err := netip.ParsePrefix(network.NetV6.String()) if err != nil { - return fmt.Errorf("parse IPv6 prefix: %w", err) + return ipv6Reassignment{}, fmt.Errorf("parse IPv6 prefix: %w", err) } - if err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix); err != nil { - return err + changed, err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix) + if err != nil { + return ipv6Reassignment{}, err } - log.WithContext(ctx).Infof("updated IPv6 addresses for %d peers in account %s (groups=%d)", - len(peers), accountID, len(settings.IPv6EnabledGroups)) + result := ipv6Reassignment{changed: changed} + for _, peer := range peers { + if peer.IPv6.IsValid() { + result.withIPv6 = append(result.withIPv6, peer.ID) + } + } - return nil + log.WithContext(ctx).Infof("updated IPv6 addresses for %d of %d peers in account %s (groups=%d)", + len(changed), len(peers), accountID, len(settings.IPv6EnabledGroups)) + + return result, nil } // reconcileIPv6ForGroupChanges checks whether the given group IDs overlap with // the account's IPv6EnabledGroups. If they do, it runs a full IPv6 address // reconciliation so that peers gaining or losing membership in an IPv6-enabled -// group get their addresses assigned or removed. -func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) error { +// group get their addresses assigned or removed. It returns the peers whose IPv6 +// address changed, which callers pass as changed peers so every peer that can +// reach them refreshes. +func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) ([]string, error) { settings, err := transaction.GetAccountSettings(ctx, store.LockingStrengthNone, accountID) if err != nil { - return fmt.Errorf("get account settings: %w", err) + return nil, fmt.Errorf("get account settings: %w", err) } if !ipv6ReconcileNeeded(settings, groupIDs) { - return nil + return nil, nil } - return am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings) + result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings) + if err != nil { + return nil, err + } + return result.changed, nil } // ipv6ReconcileNeeded reports whether changes to the given groups trigger an IPv6 @@ -2495,7 +2585,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( ctx context.Context, transaction store.Store, accountID string, peers []*nbpeer.Peer, network *types.Network, allowedPeers map[string]struct{}, v6Prefix netip.Prefix, -) error { +) ([]string, error) { takenV6 := make(map[netip.Addr]struct{}) for _, peer := range peers { if _, ok := allowedPeers[peer.ID]; ok && peer.IPv6.IsValid() && network.NetV6.Contains(peer.IPv6.AsSlice()) { @@ -2503,6 +2593,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( } } + var changed []string for _, peer := range peers { _, allowed := allowedPeers[peer.ID] oldIPv6 := peer.IPv6 @@ -2512,7 +2603,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( } else if !peer.IPv6.IsValid() || !network.NetV6.Contains(peer.IPv6.AsSlice()) { newIP, err := allocateIPv6WithRetry(v6Prefix, takenV6, peer.ID) if err != nil { - return err + return nil, err } peer.IPv6 = newIP } @@ -2522,10 +2613,11 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( } if err := transaction.SavePeer(ctx, accountID, peer); err != nil { - return fmt.Errorf("save peer %s: %w", peer.ID, err) + return nil, fmt.Errorf("save peer %s: %w", peer.ID, err) } + changed = append(changed, peer.ID) } - return nil + return changed, nil } func allocateIPv6WithRetry(prefix netip.Prefix, taken map[netip.Addr]struct{}, peerID string) (netip.Addr, error) { @@ -2569,7 +2661,7 @@ func (am *DefaultAccountManager) buildIPv6AllowedPeers(ctx context.Context, tran // Embedded proxy peers sit outside regular group membership but must // participate in any v6-enabled overlay to reach v6-only peers. - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") if err != nil { return nil, fmt.Errorf("get peers: %w", err) } @@ -2640,7 +2732,7 @@ func (am *DefaultAccountManager) updatePeerIPInTransaction(ctx context.Context, return nil } - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "") if err != nil { return fmt.Errorf("get account peers: %w", err) } diff --git a/management/server/account/deletion_hook.go b/management/server/account/deletion_hook.go new file mode 100644 index 000000000..17c817444 --- /dev/null +++ b/management/server/account/deletion_hook.go @@ -0,0 +1,14 @@ +package account + +import "context" + +// DeletionHook runs when an account is deleted, after the caller's permission to delete +// it has been checked and before any of its users or data are removed. It lets code that +// keeps per-account state outside the store tear that state down while the account still +// exists. +// +// A hook that returns an error aborts the deletion and the account is kept. The caller +// sees the error, so a hook that wants a specific response returns a status error. A +// retried deletion runs every hook again, and a later step can still fail after the hooks +// succeed, so a hook must be idempotent and must tolerate the account surviving it. +type DeletionHook func(ctx context.Context, accountID string) error diff --git a/management/server/account/manager.go b/management/server/account/manager.go index 154c9ab18..2ac8584f4 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -62,7 +62,7 @@ type Manager interface { GetUserByID(ctx context.Context, id string) (*types.User, error) GetUserFromUserAuth(ctx context.Context, userAuth auth.UserAuth) (*types.User, error) ListUsers(ctx context.Context, accountID string) ([]*types.User, error) - GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) + GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) MarkPeerConnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error MarkPeerDisconnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error DeletePeer(ctx context.Context, accountID, peerID, userID string) error diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index f31f63d0e..60075b169 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -982,18 +982,18 @@ func (mr *MockManagerMockRecorder) GetPeerNetwork(ctx, peerID any) *gomock.Call } // GetPeers mocks base method. -func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*peer.Peer, error) { +func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter) + ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter, macFilter) ret0, _ := ret[0].([]*peer.Peer) ret1, _ := ret[1].(error) return ret0, ret1 } // GetPeers indicates an expected call of GetPeers. -func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter any) *gomock.Call { +func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter, macFilter any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter, macFilter) } // GetPolicy mocks base method. diff --git a/management/server/account_test.go b/management/server/account_test.go index c63782ca8..881ad19d7 100644 --- a/management/server/account_test.go +++ b/management/server/account_test.go @@ -958,6 +958,101 @@ func TestAccountManager_DeleteAccount(t *testing.T) { assert.Len(t, pats, 0) } +func TestAccountManager_DeleteAccount_RunsDeletionHooks(t *testing.T) { + manager, _, err := createManager(t) + require.NoError(t, err) + + ownerID := "account_creator" + account, err := createAccount(manager, "test_account", ownerID, "") + require.NoError(t, err) + + // Each hook records its call and checks the account is still in the store, which is + // the point of running before deletion: a hook must be able to read what it cleans up. + var calls []string + hook := func(name string) nbAccount.DeletionHook { + return func(ctx context.Context, accountID string) error { + calls = append(calls, name+":"+accountID) + _, err := manager.Store.GetAccount(ctx, accountID) + assert.NoError(t, err, "account should still exist while hook %s runs", name) + return nil + } + } + manager.AddAccountDeletionHook(hook("first")) + manager.AddAccountDeletionHook(hook("second")) + + require.NoError(t, manager.DeleteAccount(context.Background(), account.Id, ownerID)) + + assert.Equal(t, []string{"first:" + account.Id, "second:" + account.Id}, calls, + "hooks should run once each, in registration order, with the deleted account's ID") + _, err = manager.Store.GetAccount(context.Background(), account.Id) + assert.Error(t, err, "account should be deleted after the hooks succeed") +} + +func TestAccountManager_DeleteAccount_DeletionHookErrorAbortsDeletion(t *testing.T) { + manager, _, err := createManager(t) + require.NoError(t, err) + + ownerID := "account_creator" + account, err := createAccount(manager, "test_account", ownerID, "") + require.NoError(t, err) + + manager.AddAccountDeletionHook(func(context.Context, string) error { + return status.Errorf(status.PreconditionFailed, "teardown refused") + }) + secondCalled := false + manager.AddAccountDeletionHook(func(context.Context, string) error { + secondCalled = true + return nil + }) + + err = manager.DeleteAccount(context.Background(), account.Id, ownerID) + require.Error(t, err) + + // The hook's status type has to survive the wrapping, since the HTTP layer maps it + // to the response code. + sErr, ok := status.FromError(err) + require.True(t, ok, "error should carry the hook's status error, got %v", err) + assert.Equal(t, status.PreconditionFailed, sErr.Type(), "status type should be the hook's") + assert.False(t, secondCalled, "hooks after a failing one should not run") + + _, err = manager.Store.GetAccount(context.Background(), account.Id) + assert.NoError(t, err, "account should survive a failing hook") + _, err = manager.Store.GetUserByUserID(context.Background(), store.LockingStrengthNone, ownerID) + assert.NoError(t, err, "account owner should survive a failing hook") +} + +func TestAccountManager_AddAccountDeletionHook_RejectsNil(t *testing.T) { + manager, _, err := createManager(t) + require.NoError(t, err) + + assert.PanicsWithValue(t, "nil account deletion hook", func() { + manager.AddAccountDeletionHook(nil) + }, "registering a nil hook should panic instead of breaking a later deletion") +} + +func TestAccountManager_DeleteAccount_DeletionHooksSkippedWithoutPermission(t *testing.T) { + manager, _, err := createManager(t) + require.NoError(t, err) + + ownerID := "account_creator" + account, err := createAccount(manager, "test_account", ownerID, "") + require.NoError(t, err) + + adminID := "regular_admin" + account.Users[adminID] = types.NewAdminUser(adminID) + require.NoError(t, manager.Store.SaveAccount(context.Background(), account)) + + called := false + manager.AddAccountDeletionHook(func(context.Context, string) error { + called = true + return nil + }) + + err = manager.DeleteAccount(context.Background(), account.Id, adminID) + require.Error(t, err, "only the owner may delete the account") + assert.False(t, called, "hooks should not run for a caller who may not delete the account") +} + func BenchmarkTest_GetAccountWithclaims(b *testing.B) { claims := auth.UserAuth{ Domain: "example.com", @@ -2462,7 +2557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_PeerApproval(t *testing.T) _, err = manager.UpdateAccountSettings(ctx, accountID, userID, newSettings) require.NoError(t, err) - accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, peer := range accountPeers { @@ -4462,7 +4557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) require.Len(t, peers, len(before)) for _, p := range peers { @@ -4480,7 +4575,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) for _, p := range peers { assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change for host-bit-set equivalent range", p.ID) @@ -4494,7 +4589,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) for _, p := range peers { assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change when NetworkRange omitted", p.ID) @@ -4510,7 +4605,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) for _, p := range peers { assert.True(t, newRange.Contains(p.IP), "peer %s should be in new range %s, got %s", p.ID, newRange, p.IP) @@ -4528,7 +4623,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin require.NoError(t, err) require.NotEmpty(t, settings.IPv6EnabledGroups, "new account should have IPv6 enabled for All group") - peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, p := range peers { assert.True(t, p.IPv6.IsValid(), "peer %s should have IPv6 with All group enabled", p.ID) @@ -4556,7 +4651,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin assert.Equal(t, []string{partialGroup.ID}, updatedSettings.IPv6EnabledGroups) // peer1 and peer2 should have IPv6; peer3 should not. - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) peerMap := make(map[string]*nbpeer.Peer, len(peers)) for _, p := range peers { @@ -4576,7 +4671,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin require.NoError(t, err) assert.Empty(t, updatedSettings.IPv6EnabledGroups) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, p := range peers { assert.False(t, p.IPv6.IsValid(), "peer %s should have no IPv6 when groups cleared", p.ID) @@ -4591,7 +4686,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) peerMap = make(map[string]*nbpeer.Peer, len(peers)) for _, p := range peers { diff --git a/management/server/affected_peers_ipv6_test.go b/management/server/affected_peers_ipv6_test.go new file mode 100644 index 000000000..c64360016 --- /dev/null +++ b/management/server/affected_peers_ipv6_test.go @@ -0,0 +1,243 @@ +package server + +import ( + "context" + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/controllers/network_map" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +const ( + ipv6GroupA = "ipv6-grp-a" + ipv6GroupB = "ipv6-grp-b" + ipv6GroupC = "ipv6-grp-c" + ipv6GroupD = "ipv6-grp-d" +) + +// ipv6AffectedTest holds three peers: peer1 in group A, peer2 in group B, peer3 in +// group C, with a single A<->B policy. peer3 is unrelated to peer1 and peer2. Group D +// is empty and referenced by nothing. +type ipv6AffectedTest struct { + manager *DefaultAccountManager + accountID string + peer1, peer2, peer3 *nbpeer.Peer + updMsg1, updMsg2, updMsg3 <-chan *network_map.UpdateMessage +} + +func setupIPv6AffectedTest(t *testing.T, ipv6Groups []string) *ipv6AffectedTest { + t.Helper() + + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID)) + } + + for _, g := range []*types.Group{ + {ID: ipv6GroupA, Name: "IPv6-A", Peers: []string{peer1.ID}}, + {ID: ipv6GroupB, Name: "IPv6-B", Peers: []string{peer2.ID}}, + {ID: ipv6GroupC, Name: "IPv6-C", Peers: []string{peer3.ID}}, + {ID: ipv6GroupD, Name: "IPv6-D"}, + } { + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, g)) + } + + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{{ + Enabled: true, + Sources: []string{ipv6GroupA}, + Destinations: []string{ipv6GroupB}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }}, + }, true) + require.NoError(t, err) + + // New accounts enable IPv6 for the All group; start from the requested groups. + updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = ipv6Groups + }) + + tc := &ipv6AffectedTest{ + manager: manager, + accountID: accountID, + peer1: peer1, + peer2: peer2, + peer3: peer3, + } + tc.updMsg1 = updateManager.CreateChannel(ctx, peer1.ID) + tc.updMsg2 = updateManager.CreateChannel(ctx, peer2.ID) + tc.updMsg3 = updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // The setup changes above dispatch asynchronously and can land after the + // channels open, so drop them before the test acts. + drainPeerUpdates(tc.updMsg1) + drainPeerUpdates(tc.updMsg2) + drainPeerUpdates(tc.updMsg3) + + return tc +} + +// updateIPv6TestSettings applies mutate to a copy of the current settings, so only +// the mutated fields differ from what is stored. +func updateIPv6TestSettings(t *testing.T, manager *DefaultAccountManager, accountID string, mutate func(*types.Settings)) { + t.Helper() + ctx := context.Background() + + current, err := manager.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + + updated := current.Copy() + mutate(updated) + + _, err = manager.UpdateAccountSettings(ctx, accountID, userID, updated) + require.NoError(t, err) +} + +func (tc *ipv6AffectedTest) peerIPv6(t *testing.T, peerID string) netip.Addr { + t.Helper() + peer, err := tc.manager.Store.GetPeerByID(context.Background(), store.LockingStrengthNone, tc.accountID, peerID) + require.NoError(t, err) + return peer.IPv6 +} + +func TestAffectedPeers_IPv6GroupEnabled_RefreshesOnlyReachablePeers(t *testing.T) { + tc := setupIPv6AffectedTest(t, nil) + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{ipv6GroupA} + }) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_IPv6GroupDisabled_RefreshesOnlyReachablePeers(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupA}) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should start with an IPv6 address") + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{} + }) + require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +// Widening the IPv6 range keeps peer addresses, but each holder's interface prefix +// comes from the range, so holders refresh while peers that only reach them do not. +func TestAffectedPeers_IPv6RangeWidened_RefreshesAddressHolders(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupA}) + oldIPv6 := tc.peerIPv6(t, tc.peer1.ID) + require.True(t, oldIPv6.IsValid(), "peer1 should start with an IPv6 address") + + // The range is allocated on the account network; settings may leave it empty. + network, err := tc.manager.Store.GetAccountNetwork(context.Background(), store.LockingStrengthNone, tc.accountID) + require.NoError(t, err) + current := prefixFromIPNet(network.NetV6) + require.True(t, current.IsValid(), "account should have an IPv6 range") + widened := netip.PrefixFrom(current.Addr(), current.Bits()-8).Masked() + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.NetworkRangeV6 = widened + }) + require.Equal(t, oldIPv6, tc.peerIPv6(t, tc.peer1.ID), "peer1 should keep its address inside the widened range") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldNotReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_IPv4RangeChange_RefreshesWholeAccount(t *testing.T) { + tc := setupIPv6AffectedTest(t, nil) + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.NetworkRange = netip.MustParsePrefix("100.70.0.0/16") + }) + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_IPv6WithAccountWideChange_RefreshesWholeAccount(t *testing.T) { + tc := setupIPv6AffectedTest(t, nil) + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{ipv6GroupA} + s.LazyConnectionEnabled = !s.LazyConnectionEnabled + }) + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldReceiveUpdate(t, tc.updMsg3) +} + +// Joining an IPv6-enabled group that no policy references gives peer1 an address. +// peer2 reaches peer1 through group A, not through the joined group, and must still +// learn the new address. +func TestAffectedPeers_GroupAddPeerIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupD}) + + require.NoError(t, tc.manager.GroupAddPeer(context.Background(), tc.accountID, ipv6GroupD, tc.peer1.ID)) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_UpdateGroupIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupD}) + + require.NoError(t, tc.manager.UpdateGroup(context.Background(), tc.accountID, userID, &types.Group{ + ID: ipv6GroupD, + Name: "IPv6-D", + Peers: []string{tc.peer1.ID}, + })) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +// Deleting an IPv6-enabled group removes its members' addresses after the +// pre-delete snapshot was taken. +func TestAffectedPeers_DeleteIPv6Group_RefreshesFormerMembersAndReachablePeers(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupD}) + ctx := context.Background() + + require.NoError(t, tc.manager.GroupAddPeer(ctx, tc.accountID, ipv6GroupD, tc.peer1.ID)) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + drainPeerUpdates(tc.updMsg1) + drainPeerUpdates(tc.updMsg2) + drainPeerUpdates(tc.updMsg3) + + require.NoError(t, tc.manager.DeleteGroup(ctx, tc.accountID, userID, ipv6GroupD)) + require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} diff --git a/management/server/affected_peers_user_test.go b/management/server/affected_peers_user_test.go index c0dbbb84f..3d73bbed0 100644 --- a/management/server/affected_peers_user_test.go +++ b/management/server/affected_peers_user_test.go @@ -108,11 +108,13 @@ func TestAffectedPeers_SaveUser_OnlyAffectedPeersUpdated(t *testing.T) { }) t.Run("auto group change reassigning IPv6 refreshes the changed peers and their observers", func(t *testing.T) { - account, err := manager.Store.GetAccount(ctx, accountID) - require.NoError(t, err) - account.Settings.IPv6EnabledGroups = []string{"ug-v6"} - require.NoError(t, manager.Store.SaveAccount(ctx, account)) require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-v6", Name: "ug-v6"})) + // Apply through the settings API so the reconciliation that strips the other + // peers' addresses happens here, leaving the target as the only peer the + // user update reassigns. + updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{"ug-v6"} + }) drainPeerUpdates(updTarget) drainPeerUpdates(upd2) diff --git a/management/server/affected_peers_zone_test.go b/management/server/affected_peers_zone_test.go new file mode 100644 index 000000000..4d622325c --- /dev/null +++ b/management/server/affected_peers_zone_test.go @@ -0,0 +1,145 @@ +package server + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/controllers/network_map" + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/server/affectedpeers" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +const affectedZoneDomain = "zone.test" + +// createAffectedZone stores a zone distributed to the given groups, optionally with +// one A record so the network map actually ships it. +func createAffectedZone(t *testing.T, s store.Store, accountID, domain string, enabled, withRecord bool, groups []string) *zones.Zone { + t.Helper() + ctx := context.Background() + + zone := zones.NewZone(accountID, domain, domain, enabled, false, groups) + require.NoError(t, s.CreateZone(ctx, zone)) + + if withRecord { + record := records.NewRecord(accountID, zone.ID, "host."+domain, records.RecordTypeA, "10.0.0.1", 300) + require.NoError(t, s.CreateDNSRecord(ctx, record)) + } + + return zone +} + +func TestCollectGroupChange_ZoneLinked(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]}) + + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0], "group distributed a zone should be affected by its own change") + + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Empty(t, groups, "group not referenced by any zone should not be affected") +} + +func TestCollectGroupChange_UnshippedZoneNotLinked(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Disabled zone and zone without records are never shipped by the network map. + createAffectedZone(t, s, accountID, "disabled."+affectedZoneDomain, false, true, []string{groupIDs[0]}) + createAffectedZone(t, s, accountID, "empty."+affectedZoneDomain, true, false, []string{groupIDs[1]}) + + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]}) + assert.Empty(t, groups, "groups referenced only by unshipped zones should not be affected") +} + +func TestResolveAffectedPeers_ZoneGroupMembershipChange(t *testing.T) { + _, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + + createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]}) + + // Same change shape UpdateGroup builds: the group changed as a whole and peer1 + // left it, so peer1 must refresh to drop the zone. + change := affectedpeers.Change{ + ChangedGroupIDs: []string{groupIDs[0]}, + RemovedPeersByGroup: map[string][]string{groupIDs[0]: {peerIDs[1]}}, + } + + result := resolveAffected(t, s, accountID, change) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result, "current and removed members of the zone group should be affected") +} + +func TestResolveAffectedPeers_ZoneDistributionChange(t *testing.T) { + _, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + + // Zone create/update/delete passes old and new distribution groups. + change := affectedpeers.Change{DistributionGroupIDs: []string{groupIDs[0], groupIDs[2]}} + + result := resolveAffected(t, s, accountID, change) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result, "only members of the distribution groups should be affected") +} + +// TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone verifies that adding a peer +// to a group referenced only by a zone pushes the zone to the new member and leaves +// unrelated peers alone. +func TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID)) + } + + zoneGroup := &types.Group{ID: "zone-grp", Name: "ZoneGroup", Peers: []string{peer1.ID}} + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, zoneGroup)) + + createAffectedZone(t, manager.Store, accountID, affectedZoneDomain, true, true, []string{zoneGroup.ID}) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + zoneGroup.Peers = []string{peer1.ID, peer2.ID} + require.NoError(t, manager.UpdateGroup(ctx, accountID, userID, zoneGroup)) + + peerShouldReceiveUpdate(t, updMsg1) + msg := receivePeerUpdate(t, updMsg2) + assert.True(t, syncHasCustomZone(msg, affectedZoneDomain+"."), "new zone group member should receive the zone") + peerShouldNotReceiveUpdate(t, updMsg3) +} + +func receivePeerUpdate(t *testing.T, ch <-chan *network_map.UpdateMessage) *network_map.UpdateMessage { + t.Helper() + select { + case msg := <-ch: + require.NotNil(t, msg, "update message should not be nil") + return msg + case <-time.After(peerUpdateTimeout): + require.FailNow(t, "timed out waiting for update message") + return nil + } +} + +func syncHasCustomZone(msg *network_map.UpdateMessage, domain string) bool { + for _, zone := range msg.Update.GetNetworkMap().GetDNSConfig().GetCustomZones() { + if zone.GetDomain() == domain { + return true + } + } + return false +} diff --git a/management/server/affectedpeers/resolver.go b/management/server/affectedpeers/resolver.go index cb2063ac9..895e4fd36 100644 --- a/management/server/affectedpeers/resolver.go +++ b/management/server/affectedpeers/resolver.go @@ -22,6 +22,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/internals/modules/zones" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" @@ -50,6 +51,7 @@ type Snapshot struct { policies []*types.Policy routes []*route.Route nsGroups []*nbdns.NameServerGroup + zones []*zones.Zone dnsSettings *types.DNSSettings routers []*routerTypes.NetworkRouter resources []*resourceTypes.NetworkResource @@ -127,12 +129,15 @@ func (snap *Snapshot) loadRoutesAndProxy(ctx context.Context, s store.Store, acc return snap.loadProxyServices(ctx, s, accountID) } -// loadDNS loads the nameserver groups and account DNS settings. +// loadDNS loads the nameserver groups, custom DNS zones and account DNS settings. func (snap *Snapshot) loadDNS(ctx context.Context, s store.Store, accountID string) error { var err error if snap.nsGroups, err = s.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID); err != nil { return err } + if snap.zones, err = s.GetAccountZones(ctx, store.LockingStrengthNone, accountID); err != nil { + return err + } snap.dnsSettings, err = s.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) return err } @@ -357,7 +362,7 @@ func (s policySide) opposite() policySide { // - a changed router/resource/network sits on a NETWORK -> fold the SOURCE side of // the policies whose destination reaches it (and the routers it implies). // -// Routes, nameserver groups, DNS and embedded-proxy services distribute to their own +// Routes, nameserver groups, DNS zones, DNS and embedded-proxy services distribute to their own // member peers, outside the policy graph, and are folded here too. func (r *resolver) walk() { for _, policy := range r.bothSidesPolicies() { @@ -369,6 +374,7 @@ func (r *resolver) walk() { r.collectFromPolicies() r.collectFromRoutes() r.collectFromNameServers() + r.collectFromZones() r.collectFromDNSSettings() r.collectFromNetworkRouters() r.collectFromProxyServices() @@ -829,6 +835,25 @@ func (r *resolver) collectFromNameServers() { } } +// collectFromZones folds the distribution groups of the custom DNS zones that +// reference a linked group. Like nameserver groups, a zone has no opposite side, so +// only a whole-group change folds its groups. Zones the network map does not ship +// (disabled or without records) are skipped. +func (r *resolver) collectFromZones() { + if len(r.linkGroups) == 0 { + return + } + for _, zone := range r.snap.zones { + if !zone.Enabled || len(zone.Records) == 0 { + continue + } + if anyInSet(zone.DistributionGroups, r.linkGroups) { + log.WithContext(r.ctx).Tracef("collectFromZones: zone %s references a linked group -> folding its groups %v (outputGroups only)", zone.ID, zone.DistributionGroups) + r.foldOutputGroups(zone.DistributionGroups) + } + } +} + // collectFromSSHAuthorizedGroups folds the destinations of the enabled SSH rules that // authorize a group whose user membership changed. Those destination peers carry the // group -> user mapping for the groups they authorize, so they refresh even when no diff --git a/management/server/auth/session.go b/management/server/auth/session.go index 778146589..2f1b97975 100644 --- a/management/server/auth/session.go +++ b/management/server/auth/session.go @@ -3,9 +3,11 @@ package auth import ( "context" "crypto/sha256" + "encoding/base64" "encoding/hex" "errors" "fmt" + "strings" "time" ) @@ -52,6 +54,22 @@ func (s *SessionStore) RegisterToken(ctx context.Context, token string, expiresA } func hashToken(token string) string { - sum := sha256.Sum256([]byte(token)) + sum := sha256.Sum256([]byte(canonicalizeToken(token))) return hex.EncodeToString(sum[:]) } + +// canonicalizeToken re-encodes the JWT signature segment so noncanonical +// spellings of the same signature map to one stable cache key. +func canonicalizeToken(token string) string { + parts := strings.Split(token, ".") + if len(parts) != 3 { + return token + } + + sig, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil { + return token + } + + return parts[0] + "." + parts[1] + "." + base64.RawURLEncoding.EncodeToString(sig) +} diff --git a/management/server/auth/session_test.go b/management/server/auth/session_test.go index 7c82dfc43..425e3143d 100644 --- a/management/server/auth/session_test.go +++ b/management/server/auth/session_test.go @@ -2,10 +2,15 @@ package auth import ( "context" + "crypto/rand" + "crypto/rsa" + "encoding/base64" "errors" + "strings" "testing" "time" + "github.com/golang-jwt/jwt/v5" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -131,3 +136,49 @@ func TestHashToken_StableAndDoesNotLeak(t *testing.T) { assert.Len(t, a, 64, "sha256 hex must be 64 chars") assert.NotContains(t, a, "tokenA", "raw token must not appear in hash") } + +func TestSessionStore_NoncanonicalSpellingIsRejectedAsReplay(t *testing.T) { + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{ + "sub": "user", + "exp": time.Now().Add(time.Hour).Unix(), + }) + canonical, err := token.SignedString(privateKey) + require.NoError(t, err) + + parts := strings.Split(canonical, ".") + require.Len(t, parts, 3) + + // A 256-byte RSA signature (256 mod 3 == 1) leaves unused bits in the final + // base64url character; flip one without changing the decoded signature. + const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_" + last := strings.IndexByte(alphabet, parts[2][len(parts[2])-1]) + require.GreaterOrEqual(t, last, 0) + require.Equal(t, 0, last&3, "unexpected canonical RSA signature encoding") + + equivalentSig := parts[2][:len(parts[2])-1] + string(alphabet[last|1]) + equivalent := parts[0] + "." + parts[1] + "." + equivalentSig + require.NotEqual(t, canonical, equivalent, "spellings must differ as strings") + + // Same decoded signature bytes, so they verify as the same JWT. + canonicalSig, err := base64.RawURLEncoding.DecodeString(parts[2]) + require.NoError(t, err) + altSig, err := base64.RawURLEncoding.DecodeString(equivalentSig) + require.NoError(t, err) + require.Equal(t, canonicalSig, altSig, "spellings must decode to identical signature bytes") + + // The replay-cache key must be identical for both spellings. + assert.Equal(t, hashToken(canonical), hashToken(equivalent), + "noncanonical spelling must map to the same replay-cache key") + + s := newTestSessionStore(t) + ctx := context.Background() + exp := time.Now().Add(time.Hour) + + require.NoError(t, s.RegisterToken(ctx, canonical, exp), "first claim should succeed") + err = s.RegisterToken(ctx, equivalent, exp) + require.Error(t, err, "alternate spelling must be treated as a replay") + assert.ErrorIs(t, err, ErrTokenAlreadyUsed) +} diff --git a/management/server/group.go b/management/server/group.go index 88295e2f6..8d91df3ab 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -166,9 +166,11 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use return err } - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}) + if err != nil { return err } + change.ChangedPeerIDs = ipv6Changed // A membership change does not alter which entities reference the group, so // the dependency walk runs once against the post-change snapshot. The new @@ -321,7 +323,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us var globalErr error for _, newGroup := range groups { change := affectedpeers.Change{ChangedGroupIDs: []string{newGroup.ID}} - events, snap, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change) + events, snap, change, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change) if err != nil { log.WithContext(ctx).Errorf("failed to update group %s: %v", newGroup.ID, err) if len(groups) == 1 { @@ -344,7 +346,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us return globalErr } -func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, error) { +func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, affectedpeers.Change, error) { var events []func() var snap *affectedpeers.Snapshot err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -364,9 +366,11 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI return err } - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}) + if err != nil { return err } + change.ChangedPeerIDs = ipv6Changed if err := transaction.IncrementNetworkSerial(ctx, accountID); err != nil { return err @@ -377,7 +381,7 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI snap, err = affectedpeers.Load(ctx, transaction, accountID, change) return err }) - return events, snap, err + return events, snap, change, err } // prepareGroupEvents prepares a list of event functions to be stored. @@ -480,8 +484,8 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us var allErrors error var groupIDsToDelete []string var deletedGroups []*types.Group - var snap *affectedpeers.Snapshot - var change affectedpeers.Change + var snap, ipv6Snap *affectedpeers.Snapshot + var change, ipv6Change affectedpeers.Change extraSettings, err := am.settingsManager.GetExtraSettings(ctx, accountID) if err != nil { @@ -510,10 +514,20 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us return err } - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete) + if err != nil { return err } + // Members of a deleted IPv6-enabled group lose their address, which the + // pre-delete snapshot cannot see, so they are resolved post-delete. + if len(ipv6Changed) > 0 { + ipv6Change = affectedpeers.Change{ChangedPeerIDs: ipv6Changed} + if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil { + return err + } + } + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -524,7 +538,7 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us am.StoreEvent(ctx, userID, group.ID, accountID, activity.GroupDeleted, group.EventMeta()) } - am.ExpandAndUpdateAffected(ctx, accountID, snap, change) + go am.dispatchAffected(ctx, accountID, []*affectedpeers.Snapshot{snap, ipv6Snap}, []affectedpeers.Change{change, ipv6Change}) return allErrors } @@ -564,11 +578,14 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr return err } - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}) + if err != nil { return err } + // A peer whose IPv6 address changed is visible to every peer that reaches it + // through any of its groups, not only through this one. + change.ChangedPeerIDs = ipv6Changed - var err error if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { return err } @@ -634,11 +651,14 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, return err } - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}) + if err != nil { return err } + // A peer whose IPv6 address changed is visible to every peer that reaches it + // through any of its groups, not only through this one. + change.ChangedPeerIDs = ipv6Changed - var err error if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { return err } diff --git a/management/server/http/handler.go b/management/server/http/handler.go index a57f44b3c..8ebfd6c64 100644 --- a/management/server/http/handler.go +++ b/management/server/http/handler.go @@ -58,10 +58,11 @@ import ( "github.com/netbirdio/netbird/management/server/networks/resources" "github.com/netbirdio/netbird/management/server/networks/routers" "github.com/netbirdio/netbird/management/server/telemetry" + "github.com/netbirdio/netbird/shared/ratelimit" ) // NewAPIHandler creates the Management service HTTP API handler registering all the available endpoints. -func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *middleware.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) { +func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *ratelimit.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager, proxyTokenRevocationGuard proxytoken.RevocationGuard) (http.Handler, error) { // Register bypass paths for unauthenticated endpoints if err := bypass.AddBypassPath("/api/instance"); err != nil { @@ -84,7 +85,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou if rateLimiter == nil { log.Warn("NewAPIHandler: nil rate limiter, rate limiting disabled") - rateLimiter = middleware.NewAPIRateLimiter(nil) + rateLimiter = ratelimit.NewAPIRateLimiter(nil) rateLimiter.SetEnabled(false) } @@ -135,7 +136,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou reverseproxymanager.RegisterEndpoints(serviceManager, *reverseProxyDomainManager, reverseProxyAccessLogsManager, permissionsManager, router) } - proxytoken.RegisterEndpoints(accountManager.GetStore(), permissionsManager, router) + proxytoken.RegisterEndpoints(accountManager.GetStore(), permissionsManager, proxyTokenRevocationGuard, router) // Register OAuth callback handler for proxy authentication if proxyGRPCServer != nil { diff --git a/management/server/http/handlers/accounts/accounts_handler.go b/management/server/http/handlers/accounts/accounts_handler.go index c4cba5962..795214c31 100644 --- a/management/server/http/handlers/accounts/accounts_handler.go +++ b/management/server/http/handlers/accounts/accounts_handler.go @@ -127,7 +127,7 @@ func (h *handler) validateNetworkRange(ctx context.Context, accountID, userID st } func (h *handler) validateCapacity(ctx context.Context, accountID, userID string, prefix netip.Prefix) error { - peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "") + peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "", "") if err != nil { return status.Errorf(status.Internal, "get peer count: %v", err) } diff --git a/management/server/http/handlers/groups/groups_handler.go b/management/server/http/handlers/groups/groups_handler.go index f8d161a87..1a7753a57 100644 --- a/management/server/http/handlers/groups/groups_handler.go +++ b/management/server/http/handlers/groups/groups_handler.go @@ -58,7 +58,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -77,7 +77,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -148,13 +148,10 @@ func (h *handler) updateGroup(w http.ResponseWriter, r *http.Request) { peers = *req.Peers } - resources := make([]types.Resource, 0) - if req.Resources != nil { - for _, res := range *req.Resources { - resource := types.Resource{} - resource.FromAPIRequest(&res) - resources = append(resources, resource) - } + resources, err := resourcesFromAPIRequest(req.Resources) + if err != nil { + util.WriteError(r.Context(), err, w) + return } group := types.Group{ @@ -172,7 +169,7 @@ func (h *handler) updateGroup(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -210,13 +207,10 @@ func (h *handler) createGroup(w http.ResponseWriter, r *http.Request) { peers = *req.Peers } - resources := make([]types.Resource, 0) - if req.Resources != nil { - for _, res := range *req.Resources { - resource := types.Resource{} - resource.FromAPIRequest(&res) - resources = append(resources, resource) - } + resources, err := resourcesFromAPIRequest(req.Resources) + if err != nil { + util.WriteError(r.Context(), err, w) + return } group := types.Group{ @@ -232,7 +226,7 @@ func (h *handler) createGroup(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -293,7 +287,7 @@ func (h *handler) getGroup(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -335,11 +329,30 @@ func toGroupResponse(peers []*nbpeer.Peer, group *types.Group) *api.Group { gr.PeersCount = len(gr.Peers) for _, res := range group.Resources { - resResp := res.ToAPIResponse() - gr.Resources = append(gr.Resources, *resResp) + if resResp := res.ToAPIResponse(); resResp != nil { + gr.Resources = append(gr.Resources, *resResp) + } } gr.ResourcesCount = len(gr.Resources) return &gr } + +func resourcesFromAPIRequest(req *[]api.Resource) ([]types.Resource, error) { + resources := make([]types.Resource, 0) + if req == nil { + return resources, nil + } + + for _, res := range *req { + if res.Id == "" || !types.ResourceType(res.Type).Valid() { + return nil, status.Errorf(status.InvalidArgument, "resource id shouldn't be empty and type must be one of: peer, domain, host, subnet") + } + resource := types.Resource{} + resource.FromAPIRequest(&res) + resources = append(resources, resource) + } + + return resources, nil +} diff --git a/management/server/http/handlers/groups/groups_handler_test.go b/management/server/http/handlers/groups/groups_handler_test.go index 57e238630..3e322db4e 100644 --- a/management/server/http/handlers/groups/groups_handler_test.go +++ b/management/server/http/handlers/groups/groups_handler_test.go @@ -8,8 +8,8 @@ import ( "fmt" "io" "net/http" - "net/netip" "net/http/httptest" + "net/netip" "strings" "testing" @@ -78,7 +78,7 @@ func initGroupTestData(initGroups ...*types.Group) *handler { return nil, status.Errorf(status.NotFound, "unknown group name") }, - GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { + GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { return maps.Values(TestPeers), nil }, DeleteGroupFunc: func(_ context.Context, accountID, userId, groupID string) error { @@ -208,6 +208,33 @@ func TestWriteGroup(t *testing.T) { expectedStatus: http.StatusUnprocessableEntity, expectedBody: false, }, + { + name: "Write Group POST Empty Resource", + requestType: http.MethodPost, + requestPath: "/api/groups", + requestBody: bytes.NewBuffer( + []byte(`{"name":"With Resource","resources":[{}]}`)), + expectedStatus: http.StatusUnprocessableEntity, + expectedBody: false, + }, + { + name: "Write Group PUT Empty Resource", + requestType: http.MethodPut, + requestPath: "/api/groups/id-existed", + requestBody: bytes.NewBuffer( + []byte(`{"name":"With Resource","resources":[{"id":"","type":"host"}]}`)), + expectedStatus: http.StatusUnprocessableEntity, + expectedBody: false, + }, + { + name: "Write Group POST Unknown Resource Type", + requestType: http.MethodPost, + requestPath: "/api/groups", + requestBody: bytes.NewBuffer( + []byte(`{"name":"With Resource","resources":[{"id":"res-1","type":"banana"}]}`)), + expectedStatus: http.StatusUnprocessableEntity, + expectedBody: false, + }, { name: "Write Group PUT OK", requestType: http.MethodPut, @@ -376,6 +403,20 @@ func TestGetAllGroups(t *testing.T) { } } +func TestToGroupResponseSkipsEmptyResource(t *testing.T) { + group := &types.Group{ + ID: "id-resources", + Name: "Resources", + Issued: types.GroupIssuedAPI, + Resources: []types.Resource{{}, {ID: "res-1", Type: types.ResourceTypeHost}}, + } + + got := toGroupResponse(nil, group) + + assert.Equal(t, 1, got.ResourcesCount) + assert.Equal(t, []api.Resource{{Id: "res-1", Type: api.ResourceType(types.ResourceTypeHost)}}, got.Resources) +} + func TestDeleteGroup(t *testing.T) { tt := []struct { name string diff --git a/management/server/http/handlers/peers/peers_handler.go b/management/server/http/handlers/peers/peers_handler.go index 773b640e0..8a9bf1f70 100644 --- a/management/server/http/handlers/peers/peers_handler.go +++ b/management/server/http/handlers/peers/peers_handler.go @@ -317,10 +317,11 @@ func (h *Handler) GetAllPeers(w http.ResponseWriter, r *http.Request) { nameFilter := r.URL.Query().Get("name") ipFilter := r.URL.Query().Get("ip") + macFilter := r.URL.Query().Get("mac") accountID, userID := userAuth.AccountId, userAuth.UserId - peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter) + peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter, macFilter) if err != nil { util.WriteError(r.Context(), err, w) return @@ -571,6 +572,17 @@ func peerToAccessiblePeer(peer *nbpeer.Peer, dnsDomain string) api.AccessiblePee } } +func toNetworkAddresses(addrs []nbpeer.NetworkAddress) *[]api.NetworkAddress { + if len(addrs) == 0 { + return nil + } + out := make([]api.NetworkAddress, 0, len(addrs)) + for _, a := range addrs { + out = append(out, api.NetworkAddress{NetIp: a.NetIP.String(), Mac: a.Mac}) + } + return &out +} + func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsDomain string, approved bool, reason string) *api.Peer { osVersion := peer.Meta.OSVersion if osVersion == "" { @@ -583,6 +595,7 @@ func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsD Name: peer.Name, Ip: peer.IP.String(), Ipv6: peerIPv6String(peer), + NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses), ConnectionIp: peer.Location.ConnectionIP.String(), Connected: peer.Status.Connected, LastSeen: peer.Status.LastSeen, @@ -639,6 +652,7 @@ func toPeerListItemResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dn Name: peer.Name, Ip: peer.IP.String(), Ipv6: peerIPv6String(peer), + NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses), ConnectionIp: peer.Location.ConnectionIP.String(), Connected: peer.Status.Connected, LastSeen: peer.Status.LastSeen, diff --git a/management/server/http/handlers/peers/peers_handler_test.go b/management/server/http/handlers/peers/peers_handler_test.go index 592d64d1a..7054082cc 100644 --- a/management/server/http/handlers/peers/peers_handler_test.go +++ b/management/server/http/handlers/peers/peers_handler_test.go @@ -173,7 +173,7 @@ func initTestMetaData(t *testing.T, peers ...*nbpeer.Peer) *Handler { return nil, fmt.Errorf("user not found") } }, - GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { + GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { return peers, nil }, GetPeerGroupsFunc: func(ctx context.Context, accountID, peerID string) ([]*types.Group, error) { @@ -364,6 +364,50 @@ func TestGetPeers(t *testing.T) { } } +func TestPeerResponseNetworkAddresses(t *testing.T) { + tests := []struct { + name string + addresses []nbpeer.NetworkAddress + wantJSON string + }{ + {name: "not reported"}, + {name: "empty", addresses: []nbpeer.NetworkAddress{}}, + { + name: "multiple interfaces", + addresses: []nbpeer.NetworkAddress{ + {NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"}, + {NetIP: netip.MustParsePrefix("2001:db8::123/64"), Mac: "00:93:37:bd:83:10"}, + }, + wantJSON: `[{"net_ip":"192.168.0.11/24","mac":"00:93:37:bd:83:0f"},{"net_ip":"2001:db8::123/64","mac":"00:93:37:bd:83:10"}]`, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peer := &nbpeer.Peer{ + Status: &nbpeer.PeerStatus{}, + Meta: nbpeer.PeerSystemMeta{NetworkAddresses: tt.addresses}, + } + responses := map[string]any{ + "single peer": toSinglePeerResponse(peer, nil, "example.com", true, ""), + "peer list": toPeerListItemResponse(peer, nil, "example.com", 0), + } + for name, response := range responses { + t.Run(name, func(t *testing.T) { + body, err := json.Marshal(response) + require.NoError(t, err) + var fields map[string]json.RawMessage + require.NoError(t, json.Unmarshal(body, &fields)) + if tt.wantJSON == "" { + assert.NotContains(t, fields, "network_addresses", "unreported interfaces should be omitted") + return + } + assert.JSONEq(t, tt.wantJSON, string(fields["network_addresses"]), "response should preserve interface addresses and MACs") + }) + } + }) + } +} + func TestGetAccessiblePeers(t *testing.T) { peer1 := &nbpeer.Peer{ ID: "peer1", diff --git a/management/server/http/handlers/proxy/auth.go b/management/server/http/handlers/proxy/auth.go index 0f4b72e14..133236401 100644 --- a/management/server/http/handlers/proxy/auth.go +++ b/management/server/http/handlers/proxy/auth.go @@ -16,21 +16,21 @@ import ( "golang.org/x/oauth2" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" - "github.com/netbirdio/netbird/management/server/http/middleware" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/proxy/auth" + "github.com/netbirdio/netbird/shared/ratelimit" ) // AuthCallbackHandler handles OAuth callbacks for proxy authentication. type AuthCallbackHandler struct { proxyService *nbgrpc.ProxyServiceServer - rateLimiter *middleware.APIRateLimiter + rateLimiter *ratelimit.APIRateLimiter trustedProxies []netip.Prefix } // NewAuthCallbackHandler creates a new OAuth callback handler. func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProxies []netip.Prefix) *AuthCallbackHandler { - rateLimiterConfig := &middleware.RateLimiterConfig{ + rateLimiterConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 10, Burst: 15, CleanupInterval: 5 * time.Minute, @@ -39,7 +39,7 @@ func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProx return &AuthCallbackHandler{ proxyService: proxyService, - rateLimiter: middleware.NewAPIRateLimiter(rateLimiterConfig), + rateLimiter: ratelimit.NewAPIRateLimiter(rateLimiterConfig), trustedProxies: trustedProxies, } } @@ -59,7 +59,7 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ state := r.URL.Query().Get("state") - codeVerifier, originalURL, err := h.proxyService.ValidateState(state) + codeVerifier, originalURL, useSessionCode, err := h.proxyService.ValidateState(state) if err != nil { log.WithError(err).Error("OAuth callback state validation failed") http.Error(w, "Invalid state parameter", http.StatusBadRequest) @@ -119,10 +119,19 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ redirectURL.Scheme = "https" query := redirectURL.Query() - query.Set("session_token", sessionToken) + if useSessionCode { + code, ok := h.proxyService.GenerateSessionCode(sessionToken) + if !ok { + http.Error(w, "Failed to create session", http.StatusInternalServerError) + return + } + query.Set(auth.SessionCodeQueryParam, code) + } else { + query.Set(auth.SessionTokenQueryParam, sessionToken) + } redirectURL.RawQuery = query.Encode() - log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user with session token") + log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user to proxy") http.Redirect(w, r, redirectURL.String(), http.StatusFound) } diff --git a/management/server/http/handlers/proxy/auth_callback_integration_test.go b/management/server/http/handlers/proxy/auth_callback_integration_test.go index 1dbfca4cd..862d5d5f2 100644 --- a/management/server/http/handlers/proxy/auth_callback_integration_test.go +++ b/management/server/http/handlers/proxy/auth_callback_integration_test.go @@ -181,6 +181,10 @@ func (m *testAccessLogManager) GetAllAccessLogs(_ context.Context, _, _ string, } func setupAuthCallbackTest(t *testing.T) *testSetup { + return setupAuthCallbackTestWithProxyManager(t, testSessionCodeManager{}) +} + +func setupAuthCallbackTestWithProxyManager(t *testing.T, proxyManager nbproxy.Manager) *testSetup { t.Helper() ctx := context.Background() @@ -197,7 +201,7 @@ func setupAuthCallbackTest(t *testing.T) *testSetup { require.NoError(t, err) tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) - pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore) usersManager := users.NewManager(testStore) @@ -212,12 +216,12 @@ func setupAuthCallbackTest(t *testing.T) *testSetup { proxyService := nbgrpc.NewProxyServiceServer( &testAccessLogManager{}, tokenStore, - pkceStore, + singleUseStore, oidcConfig, nil, usersManager, nil, - nil, + proxyManager, nil, ) @@ -242,6 +246,15 @@ func setupAuthCallbackTest(t *testing.T) *testSetup { } } +type testSessionCodeManager struct { + nbproxy.Manager + supported bool +} + +func (m testSessionCodeManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool { + return m.supported +} + func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store.Store) { t.Helper() @@ -252,10 +265,11 @@ func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store privKey := base64.StdEncoding.EncodeToString(priv) testProxy := &service.Service{ - ID: "testProxyId", - AccountID: "testAccountId", - Name: "Test Proxy", - Domain: "test-proxy.example.com", + ID: "testProxyId", + AccountID: "testAccountId", + Name: "Test Proxy", + Domain: "test-proxy.example.com", + ProxyCluster: "cluster.example.com", Targets: []*service.Target{{ Path: strPtr("/"), Host: "localhost", @@ -512,29 +526,56 @@ func createTestState(t *testing.T, ps *nbgrpc.ProxyServiceServer, redirectURL st } func TestAuthCallback_UserAllowedToLogin(t *testing.T) { - setup := setupAuthCallbackTest(t) - defer setup.cleanup() + tests := []struct { + name string + manager nbproxy.Manager + wantParam string + absentParam string + }{ + {name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "nb_session_code"}, + {name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "nb_session_code", absentParam: "session_token"}, + } - setup.oidcServer.tokenSubject = "allowedUserId" + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + setup := setupAuthCallbackTestWithProxyManager(t, tt.manager) + defer setup.cleanup() - state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard") + setup.oidcServer.tokenSubject = "allowedUserId" + state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard") + req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil) + rec := httptest.NewRecorder() + setup.router.ServeHTTP(rec, req) + require.Equal(t, http.StatusFound, rec.Code) - req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil) - rec := httptest.NewRecorder() + location, err := url.Parse(rec.Header().Get("Location")) + require.NoError(t, err) + require.Equal(t, "test-proxy.example.com", location.Host) + require.NotEmpty(t, location.Query().Get(tt.wantParam)) + require.Empty(t, location.Query().Get(tt.absentParam)) + require.Empty(t, location.Query().Get("error")) - setup.router.ServeHTTP(rec, req) + if tt.wantParam == "nb_session_code" { + code := location.Query().Get("nb_session_code") + response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ + Domain: location.Hostname(), + SessionCode: code, + }) + require.NoError(t, err) + require.True(t, response.GetValid()) + require.NotEmpty(t, response.GetSessionToken()) + require.NotEqual(t, code, response.GetSessionToken()) - require.Equal(t, http.StatusFound, rec.Code) - - location := rec.Header().Get("Location") - require.NotEmpty(t, location) - - parsedLocation, err := url.Parse(location) - require.NoError(t, err) - - require.Equal(t, "test-proxy.example.com", parsedLocation.Host) - require.NotEmpty(t, parsedLocation.Query().Get("session_token"), "Should include session token") - require.Empty(t, parsedLocation.Query().Get("error"), "Should not have error parameter") + replayed, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ + Domain: location.Hostname(), + SessionCode: code, + }) + require.NoError(t, err) + require.False(t, replayed.GetValid()) + require.Empty(t, replayed.GetSessionToken()) + } + }) + } } // TestAuthCallback_UserDeniedByAccountStatus asserts that a user whose account diff --git a/management/server/http/handlers/users/invites_handler.go b/management/server/http/handlers/users/invites_handler.go index 0f0f57c29..d96e58bde 100644 --- a/management/server/http/handlers/users/invites_handler.go +++ b/management/server/http/handlers/users/invites_handler.go @@ -11,15 +11,15 @@ import ( "github.com/netbirdio/netbird/management/server/account" nbcontext "github.com/netbirdio/netbird/management/server/context" - "github.com/netbirdio/netbird/management/server/http/middleware" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/http/api" "github.com/netbirdio/netbird/shared/management/http/util" "github.com/netbirdio/netbird/shared/management/status" + "github.com/netbirdio/netbird/shared/ratelimit" ) // publicInviteRateLimiter limits public invite requests by IP address to prevent brute-force attacks -var publicInviteRateLimiter = middleware.NewAPIRateLimiter(&middleware.RateLimiterConfig{ +var publicInviteRateLimiter = ratelimit.NewAPIRateLimiter(&ratelimit.RateLimiterConfig{ RequestsPerMinute: 10, // 10 attempts per minute per IP Burst: 5, // Allow burst of 5 requests CleanupInterval: 10 * time.Minute, diff --git a/management/server/http/middleware/auth_middleware.go b/management/server/http/middleware/auth_middleware.go index ba8c66241..1e831e6c3 100644 --- a/management/server/http/middleware/auth_middleware.go +++ b/management/server/http/middleware/auth_middleware.go @@ -18,6 +18,7 @@ import ( "github.com/netbirdio/netbird/shared/auth" "github.com/netbirdio/netbird/shared/management/http/util" "github.com/netbirdio/netbird/shared/management/status" + "github.com/netbirdio/netbird/shared/ratelimit" ) type EnsureAccountFunc func(ctx context.Context, userAuth auth.UserAuth) (string, string, error) @@ -33,7 +34,7 @@ type AuthMiddleware struct { ensureAccount EnsureAccountFunc getUserFromUserAuth GetUserFromUserAuthFunc syncUserJWTGroups SyncUserJWTGroupsFunc - rateLimiter *APIRateLimiter + rateLimiter *ratelimit.APIRateLimiter patUsageTracker *PATUsageTracker isValidChildAccount IsValidChildAccountFunc } @@ -44,7 +45,7 @@ func NewAuthMiddleware( ensureAccount EnsureAccountFunc, syncUserJWTGroups SyncUserJWTGroupsFunc, getUserFromUserAuth GetUserFromUserAuthFunc, - rateLimiter *APIRateLimiter, + rateLimiter *ratelimit.APIRateLimiter, meter metric.Meter, isValidChildAccount IsValidChildAccountFunc, ) *AuthMiddleware { diff --git a/management/server/http/middleware/auth_middleware_test.go b/management/server/http/middleware/auth_middleware_test.go index a34554660..ceefc7c45 100644 --- a/management/server/http/middleware/auth_middleware_test.go +++ b/management/server/http/middleware/auth_middleware_test.go @@ -18,6 +18,7 @@ import ( "github.com/netbirdio/netbird/management/server/util" nbauth "github.com/netbirdio/netbird/shared/auth" nbjwt "github.com/netbirdio/netbird/shared/auth/jwt" + "github.com/netbirdio/netbird/shared/ratelimit" ) const ( @@ -196,7 +197,7 @@ func TestAuthMiddleware_Handler(t *testing.T) { GetPATInfoFunc: mockGetAccountInfoFromPAT, } - disabledLimiter := NewAPIRateLimiter(nil) + disabledLimiter := ratelimit.NewAPIRateLimiter(nil) disabledLimiter.SetEnabled(false) authMiddleware := NewAuthMiddleware( mockAuth, @@ -260,7 +261,7 @@ func TestAuthMiddleware_SyncUserJWTGroupsDetachedFromRequestCancellation(t *test GetPATInfoFunc: mockGetAccountInfoFromPAT, } - disabledLimiter := NewAPIRateLimiter(nil) + disabledLimiter := ratelimit.NewAPIRateLimiter(nil) disabledLimiter.SetEnabled(false) authMiddleware := NewAuthMiddleware( @@ -311,7 +312,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { t.Run("PAT Token Rate Limiting - Burst Works", func(t *testing.T) { // Configure rate limiter: 10 requests per minute with burst of 5 - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 10, Burst: 5, CleanupInterval: 5 * time.Minute, @@ -329,7 +330,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -364,7 +365,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { t.Run("PAT Token Rate Limiting - Rate Limit Enforced", func(t *testing.T) { // Configure very low rate limit: 1 request per minute - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 1, Burst: 1, CleanupInterval: 5 * time.Minute, @@ -382,7 +383,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -408,7 +409,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { t.Run("Bearer Token Not Rate Limited", func(t *testing.T) { // Configure strict rate limit - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 1, Burst: 1, CleanupInterval: 5 * time.Minute, @@ -426,7 +427,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -453,7 +454,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { t.Run("PAT Token Rate Limiting Per Token", func(t *testing.T) { // Configure rate limiter - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 1, Burst: 1, CleanupInterval: 5 * time.Minute, @@ -471,7 +472,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -518,7 +519,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { t.Run("Rate Limiter Cleanup", func(t *testing.T) { // Configure rate limiter with short cleanup interval and TTL for testing - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 60, Burst: 1, CleanupInterval: 100 * time.Millisecond, @@ -536,7 +537,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -578,7 +579,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { }) t.Run("Terraform User Agent Not Rate Limited", func(t *testing.T) { - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 1, Burst: 1, CleanupInterval: 5 * time.Minute, @@ -596,7 +597,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -634,7 +635,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { }) t.Run("Non-Terraform User Agent With PAT Is Rate Limited", func(t *testing.T) { - rateLimitConfig := &RateLimiterConfig{ + rateLimitConfig := &ratelimit.RateLimiterConfig{ RequestsPerMinute: 1, Burst: 1, CleanupInterval: 5 * time.Minute, @@ -652,7 +653,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) { func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) { return &types.User{}, nil }, - NewAPIRateLimiter(rateLimitConfig), + ratelimit.NewAPIRateLimiter(rateLimitConfig), nil, func(_ context.Context, _, _, _ string) bool { return false }, ) @@ -740,7 +741,7 @@ func TestAuthMiddleware_Handler_Child(t *testing.T) { GetPATInfoFunc: mockGetAccountInfoFromPAT, } - disabledLimiter := NewAPIRateLimiter(nil) + disabledLimiter := ratelimit.NewAPIRateLimiter(nil) disabledLimiter.SetEnabled(false) authMiddleware := NewAuthMiddleware( mockAuth, diff --git a/management/server/http/testing/testing_tools/channel/channel.go b/management/server/http/testing/testing_tools/channel/channel.go index ab992b1ca..c3f6a06e0 100644 --- a/management/server/http/testing/testing_tools/channel/channel.go +++ b/management/server/http/testing/testing_tools/channel/channel.go @@ -46,14 +46,14 @@ import ( "github.com/netbirdio/netbird/management/server/networks/routers" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/settings" - "github.com/netbirdio/netbird/management/server/store" + nbstore "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/users" "github.com/netbirdio/netbird/shared/auth" ) func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPeerUpdate *network_map.UpdateMessage, validateUpdate bool) (http.Handler, account.Manager, chan struct{}) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) + store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) if err != nil { t.Fatalf("Failed to create test store: %v", err) } @@ -108,15 +108,15 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee t.Fatalf("Failed to create manager: %v", err) } - accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil) + accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil) proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) - pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore) noopMeter := noop.NewMeterProvider().Meter("") proxyMgr, err := proxymanager.NewManager(store, noopMeter) if err != nil { t.Fatalf("Failed to create proxy manager: %v", err) } - proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, pkceverifierStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil) + proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil) // NewProxyServiceServer starts cleanupStaleProxies on a context it derives // from context.Background(), independent of the cancellable ctx above; // Close() cancels it so the goroutine does not outlive the test. @@ -147,7 +147,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager) apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter() - apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil) + apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil) if err != nil { t.Fatalf("Failed to create API handler: %v", err) } @@ -204,7 +204,7 @@ func PeerShouldNotReceiveAnyUpdate(t testing_tools.TB, updateMessage <-chan *net // BuildApiBlackBoxWithDBStateAndPeerChannel creates the API handler and returns // the peer update channel directly so tests can verify updates inline. func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile string) (http.Handler, account.Manager, <-chan *network_map.UpdateMessage) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) + store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) if err != nil { t.Fatalf("Failed to create test store: %v", err) } @@ -248,15 +248,15 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin t.Fatalf("Failed to create manager: %v", err) } - accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil) + accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil) proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) - pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore) noopMeter := noop.NewMeterProvider().Meter("") proxyMgr, err := proxymanager.NewManager(store, noopMeter) if err != nil { t.Fatalf("Failed to create proxy manager: %v", err) } - proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, pkceverifierStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil) + proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil) // NewProxyServiceServer starts cleanupStaleProxies on a context it derives // from context.Background(), independent of the cancellable ctx above; // Close() cancels it so the goroutine does not outlive the test. @@ -287,7 +287,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager) apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter() - apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil) + apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil) if err != nil { t.Fatalf("Failed to create API handler: %v", err) } diff --git a/management/server/idp/auth0.go b/management/server/idp/auth0.go index 7d3837190..4d6ef5859 100644 --- a/management/server/idp/auth0.go +++ b/management/server/idp/auth0.go @@ -132,13 +132,7 @@ type ConnectionOptions struct { // NewAuth0Manager creates a new instance of the Auth0Manager func NewAuth0Manager(config Auth0ClientConfig, appMetrics telemetry.AppMetrics) (*Auth0Manager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/authentik.go b/management/server/idp/authentik.go index ebd79b715..9ab884bc7 100644 --- a/management/server/idp/authentik.go +++ b/management/server/idp/authentik.go @@ -49,13 +49,7 @@ type AuthentikCredentials struct { // NewAuthentikManager creates a new instance of the AuthentikManager. func NewAuthentikManager(config AuthentikClientConfig, appMetrics telemetry.AppMetrics) (*AuthentikManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/azure.go b/management/server/idp/azure.go index 320ca7a83..6640ef6d0 100644 --- a/management/server/idp/azure.go +++ b/management/server/idp/azure.go @@ -54,13 +54,7 @@ type azureProfile map[string]any // NewAzureManager creates a new instance of the AzureManager. func NewAzureManager(config AzureClientConfig, appMetrics telemetry.AppMetrics) (*AzureManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/dex.go b/management/server/idp/dex.go index 0cac246e1..7d25c6ed0 100644 --- a/management/server/idp/dex.go +++ b/management/server/idp/dex.go @@ -4,10 +4,8 @@ import ( "context" "encoding/base64" "fmt" - "net/http" "strings" "sync" - "time" "github.com/dexidp/dex/api/v2" log "github.com/sirupsen/logrus" @@ -44,13 +42,7 @@ func NewDexManager(config DexClientConfig, appMetrics telemetry.AppMetrics) (*De return nil, fmt.Errorf("dex IdP configuration is incomplete, GRPCAddr is missing") } - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: 10 * time.Second, - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} return &DexManager{ diff --git a/management/server/idp/google_workspace.go b/management/server/idp/google_workspace.go index dadbfd83e..ff58e7772 100644 --- a/management/server/idp/google_workspace.go +++ b/management/server/idp/google_workspace.go @@ -4,7 +4,6 @@ import ( "context" "encoding/base64" "fmt" - "net/http" log "github.com/sirupsen/logrus" "golang.org/x/oauth2/google" @@ -44,13 +43,7 @@ func (gc *GoogleWorkspaceCredentials) Authenticate(_ context.Context) (JWTToken, // NewGoogleWorkspaceManager creates a new instance of the GoogleWorkspaceManager. func NewGoogleWorkspaceManager(ctx context.Context, config GoogleWorkspaceClientConfig, appMetrics telemetry.AppMetrics) (*GoogleWorkspaceManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/jumpcloud.go b/management/server/idp/jumpcloud.go index f0dec3a9b..ac547e4a1 100644 --- a/management/server/idp/jumpcloud.go +++ b/management/server/idp/jumpcloud.go @@ -58,13 +58,7 @@ type JumpCloudCredentials struct { // NewJumpCloudManager creates a new instance of the JumpCloudManager. func NewJumpCloudManager(config JumpCloudClientConfig, appMetrics telemetry.AppMetrics) (*JumpCloudManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/keycloak.go b/management/server/idp/keycloak.go index 1cf26394f..9c01fcee2 100644 --- a/management/server/idp/keycloak.go +++ b/management/server/idp/keycloak.go @@ -59,13 +59,7 @@ type keycloakProfile struct { // NewKeycloakManager creates a new instance of the KeycloakManager. func NewKeycloakManager(config KeycloakClientConfig, appMetrics telemetry.AppMetrics) (*KeycloakManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/okta.go b/management/server/idp/okta.go index 07f0d8008..90bcd05a9 100644 --- a/management/server/idp/okta.go +++ b/management/server/idp/okta.go @@ -40,13 +40,7 @@ type OktaCredentials struct { // NewOktaManager creates a new instance of the OktaManager. func NewOktaManager(config OktaClientConfig, appMetrics telemetry.AppMetrics) (*OktaManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} config.Issuer = baseURL(config.Issuer) diff --git a/management/server/idp/pocketid.go b/management/server/idp/pocketid.go index fc338b86b..b340bfe5f 100644 --- a/management/server/idp/pocketid.go +++ b/management/server/idp/pocketid.go @@ -83,13 +83,7 @@ type pocketIdUserGroupDto struct { } func NewPocketIdManager(config PocketIdClientConfig, appMetrics telemetry.AppMetrics) (*PocketIdManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/idp/util.go b/management/server/idp/util.go index 6545c2a69..be59edba7 100644 --- a/management/server/idp/util.go +++ b/management/server/idp/util.go @@ -2,6 +2,8 @@ package idp import ( "encoding/json" + "errors" + "net/http" "net/url" "os" "strings" @@ -81,6 +83,23 @@ const ( defaultTimeout = 10 * time.Second ) +// errRedirectRefused is returned instead of http.ErrUseLastResponse so the +// client closes the redirect response rather than handing it back unread. +var errRedirectRefused = errors.New("redirect refused") + +func newHTTPClient() *http.Client { + httpTransport := http.DefaultTransport.(*http.Transport).Clone() + httpTransport.MaxIdleConns = 5 + + return &http.Client{ + Timeout: idpTimeout(), + Transport: httpTransport, + CheckRedirect: func(*http.Request, []*http.Request) error { + return errRedirectRefused + }, + } +} + // idpTimeout returns a timeout value for the IDP func idpTimeout() time.Duration { timeoutStr, ok := os.LookupEnv(idpTimeoutEnv) diff --git a/management/server/idp/zitadel.go b/management/server/idp/zitadel.go index 320f0c131..fdc59915d 100644 --- a/management/server/idp/zitadel.go +++ b/management/server/idp/zitadel.go @@ -160,13 +160,7 @@ func verifyJWTConfig(config ZitadelClientConfig) error { // NewZitadelManager creates a new instance of the ZitadelManager. func NewZitadelManager(config ZitadelClientConfig, appMetrics telemetry.AppMetrics) (*ZitadelManager, error) { - httpTransport := http.DefaultTransport.(*http.Transport).Clone() - httpTransport.MaxIdleConns = 5 - - httpClient := &http.Client{ - Timeout: idpTimeout(), - Transport: httpTransport, - } + httpClient := newHTTPClient() helper := JsonParser{} diff --git a/management/server/integrated_validator.go b/management/server/integrated_validator.go index 9ec1f491e..5928a8ed2 100644 --- a/management/server/integrated_validator.go +++ b/management/server/integrated_validator.go @@ -100,7 +100,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI return nil, nil, err } - peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") if err != nil { return nil, nil, err } diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index 2f871c3e2..3313bf99c 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -39,7 +39,7 @@ type MockAccountManager struct { GetAccountIDByUserIdFunc func(ctx context.Context, userAuth auth.UserAuth) (string, error) GetUserFromUserAuthFunc func(ctx context.Context, userAuth auth.UserAuth) (*types.User, error) ListUsersFunc func(ctx context.Context, accountID string) ([]*types.User, error) - GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) + GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) MarkPeerConnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error MarkPeerDisconnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error SyncAndMarkPeerFunc func(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) @@ -807,9 +807,9 @@ func (am *MockAccountManager) GetAccountIDFromUserAuth(ctx context.Context, user } // GetPeers mocks GetPeers of the AccountManager interface -func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { +func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { if am.GetPeersFunc != nil { - return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter) + return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter, macFilter) } return nil, status.Errorf(codes.Unimplemented, "method GetPeers is not implemented") } diff --git a/management/server/networks/resources/types/resource.go b/management/server/networks/resources/types/resource.go index 4cf7f7ea3..bb33e00eb 100644 --- a/management/server/networks/resources/types/resource.go +++ b/management/server/networks/resources/types/resource.go @@ -32,7 +32,7 @@ type NetworkResource struct { ID string `gorm:"primaryKey"` NetworkID string `gorm:"index"` AccountID string `gorm:"index"` - PublicID string `json:"-"` + PublicID string `json:"-" gorm:"index"` Name string Description string Type NetworkResourceType diff --git a/management/server/peer.go b/management/server/peer.go index 9f5572252..5d5863fa7 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -47,7 +47,7 @@ const ( // GetPeers returns peers visible to the user within an account. // Users with "peers:read" see all peers. Otherwise, users see only their own peers, or none if restricted by account settings. -func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { +func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { user, err := am.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID) if err != nil { return nil, err @@ -59,7 +59,7 @@ func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID } if allowed { - return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter) + return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter, macFilter) } settings, err := am.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID) diff --git a/management/server/peer_test.go b/management/server/peer_test.go index 22f2b9b6f..5c3e02af5 100644 --- a/management/server/peer_test.go +++ b/management/server/peer_test.go @@ -4,10 +4,14 @@ import ( "context" "crypto/sha256" b64 "encoding/base64" + "encoding/json" "fmt" "io" "net" + "net/http" + "net/http/httptest" "net/netip" + "net/url" "os" "runtime" "strconv" @@ -33,12 +37,15 @@ import ( "github.com/netbirdio/netbird/management/internals/server/config" "github.com/netbirdio/netbird/management/internals/shared/grpc" nbcache "github.com/netbirdio/netbird/management/server/cache" + nbcontext "github.com/netbirdio/netbird/management/server/context" + peershandler "github.com/netbirdio/netbird/management/server/http/handlers/peers" "github.com/netbirdio/netbird/management/server/http/testing/testing_tools" "github.com/netbirdio/netbird/management/server/integrations/port_forwarding" "github.com/netbirdio/netbird/management/server/job" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/shared/auth" + "github.com/netbirdio/netbird/shared/management/http/api" "github.com/netbirdio/netbird/shared/management/status" "github.com/netbirdio/netbird/management/server/util" @@ -718,7 +725,7 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) { return } - peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "") + peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "", "") if err != nil { t.Fatal(err) return @@ -731,6 +738,71 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) { } } +func TestDefaultAccountManager_GetPeers_FilterByMac(t *testing.T) { + ctx := context.Background() + manager, _, err := createManager(t) + require.NoError(t, err) + account := newAccountWithId(ctx, "mac-account", "mac-admin", "", "", "", false) + account.Peers["matching"] = &nbpeer.Peer{ + ID: "matching", Key: "matching-key", Name: "laptop", DNSLabel: "laptop", + IP: netip.MustParseAddr("100.64.0.10"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()}, + Meta: nbpeer.PeerSystemMeta{NetworkAddresses: []nbpeer.NetworkAddress{ + {NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"}, + {NetIP: netip.MustParsePrefix("192.168.1.11/24"), Mac: "aa:bb:cc:dd:ee:ff"}, + }}, + } + account.Peers["other"] = &nbpeer.Peer{ + ID: "other", Key: "other-key", Name: "desktop", DNSLabel: "desktop", + IP: netip.MustParseAddr("100.64.0.20"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()}, + } + require.NoError(t, manager.Store.SaveAccount(ctx, account)) + otherAccount := newAccountWithId(ctx, "other-account", "other-admin", "", "", "", false) + otherPeer := account.Peers["matching"].Copy() + otherPeer.ID, otherPeer.Key = "outside-account", "outside-key" + otherAccount.Peers[otherPeer.ID] = otherPeer + require.NoError(t, manager.Store.SaveAccount(ctx, otherAccount)) + handler := peershandler.NewHandler(manager, manager.networkMapController, manager.permissionsManager) + + tests := []struct { + name, nameFilter, ipFilter, macFilter string + wantIDs []string + }{ + {name: "no filter", wantIDs: []string{"matching", "other"}}, + {name: "full MAC", macFilter: "00:93:37:bd:83:0f", wantIDs: []string{"matching"}}, + {name: "partial MAC", macFilter: "93:37:bd", wantIDs: []string{"matching"}}, + {name: "second interface", macFilter: "aa:bb:cc:dd:ee:ff", wantIDs: []string{"matching"}}, + {name: "unknown MAC", macFilter: "11:22:33:44:55:66"}, + {name: "combined filters", nameFilter: "laptop", ipFilter: "100.64.0.10", macFilter: "00:93:37", wantIDs: []string{"matching"}}, + {name: "name mismatch", nameFilter: "desktop", macFilter: "00:93:37"}, + {name: "IP mismatch", ipFilter: "100.64.0.20", macFilter: "00:93:37"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := manager.GetPeers(ctx, account.Id, "mac-admin", tt.nameFilter, tt.ipFilter, tt.macFilter) + require.NoError(t, err) + ids := make([]string, 0, len(peers)) + for _, peer := range peers { + ids = append(ids, peer.ID) + } + assert.ElementsMatch(t, tt.wantIDs, ids, "filters should return only matching peers in the account") + + query := url.Values{"name": {tt.nameFilter}, "ip": {tt.ipFilter}, "mac": {tt.macFilter}} + req := httptest.NewRequest(http.MethodGet, "/api/peers?"+query.Encode(), nil) + req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: account.Id, UserId: "mac-admin"}) + recorder := httptest.NewRecorder() + handler.GetAllPeers(recorder, req) + require.Equal(t, http.StatusOK, recorder.Code, "peer listing should succeed: %s", recorder.Body.String()) + var response []api.PeerBatch + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) + responseIDs := make([]string, 0, len(response)) + for _, peer := range response { + responseIDs = append(responseIDs, peer.Id) + } + assert.ElementsMatch(t, tt.wantIDs, responseIDs, "HTTP query filters should reach the store") + }) + } +} + func setupTestAccountManager(b testing.TB, peers int, groups int) (*DefaultAccountManager, *update_channel.PeersUpdateManager, string, string, error) { b.Helper() @@ -934,7 +1006,7 @@ func BenchmarkGetPeers(b *testing.B) { b.ResetTimer() for i := 0; i < b.N; i++ { - _, err := manager.GetPeers(context.Background(), accountID, userID, "", "") + _, err := manager.GetPeers(context.Background(), accountID, userID, "", "", "") if err != nil { b.Fatalf("GetPeers failed: %v", err) } diff --git a/management/server/permissions/agent_network_roles_test.go b/management/server/permissions/agent_network_roles_test.go index 9ab708bd7..f5ad2000d 100644 --- a/management/server/permissions/agent_network_roles_test.go +++ b/management/server/permissions/agent_network_roles_test.go @@ -62,11 +62,11 @@ func TestAgentNetworkAdminRole(t *testing.T) { } } -// TestUsageViewerRole pins the least-privilege cost role: read on the -// aggregated usage overview plus read-only on the resources its filters -// and display columns resolve against (users, groups, peers, the provider -// list) — no policies, no request-level logs (which can contain captured -// prompts), nothing else in the account. +// TestUsageViewerRole pins the read-only usage role: read on the aggregated +// usage overview and the account-wide request-level logs, plus read-only on +// the resources their filters and display columns resolve against (users, +// groups, peers, the provider list) — no policies, guardrails, budgets, or +// settings, nothing else in the account. func TestUsageViewerRole(t *testing.T) { manager := NewManager(nil) ctx := context.Background() @@ -76,6 +76,7 @@ func TestUsageViewerRole(t *testing.T) { readOnly := []modules.Module{ modules.AgentNetworkUsage, + modules.AgentNetworkLogs, modules.AgentNetworkProviders, modules.Users, modules.Groups, @@ -83,7 +84,7 @@ func TestUsageViewerRole(t *testing.T) { } for _, m := range readOnly { assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, m, operations.Read), - "usage_viewer must read %s for the usage view and its filters", m) + "usage_viewer must read %s for the usage and log views and their filters", m) for _, op := range []operations.Operation{operations.Create, operations.Update, operations.Delete} { assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, m, op), "usage_viewer must not have %s on %s", op, m) @@ -95,7 +96,6 @@ func TestUsageViewerRole(t *testing.T) { modules.AgentNetworkPolicies, modules.AgentNetworkGuardrails, modules.AgentNetworkBudgets, - modules.AgentNetworkLogs, modules.AgentNetworkSettings, modules.Networks, modules.SetupKeys, diff --git a/management/server/permissions/roles/usage_viewer.go b/management/server/permissions/roles/usage_viewer.go index e480ae478..ab35a24db 100644 --- a/management/server/permissions/roles/usage_viewer.go +++ b/management/server/permissions/roles/usage_viewer.go @@ -7,16 +7,15 @@ import ( ) // UsageViewer is the regular User baseline plus read access to the -// aggregated Agent Network usage and cost overview, and read-only access -// to the resources the usage filters and display columns resolve against: -// users and groups (identity filters and name resolution), peers (agent -// principals in the caller column), and the provider list (provider and -// model filter options — the manager redacts connection config such as -// upstream URLs and operator-supplied header values for callers holding -// read without update). It sees no policies and no account-wide -// request-level access logs (which can contain captured prompts); its own -// requests remain readable through the self-scoped endpoints, like any -// caller's. +// aggregated Agent Network usage and cost overview and to the account-wide +// request-level access logs (which can contain captured prompts), and +// read-only access to the resources the usage and log filters and display +// columns resolve against: users and groups (identity filters and name +// resolution), peers (agent principals in the caller column), and the +// provider list (provider and model filter options — the manager redacts +// connection config such as upstream URLs and operator-supplied header +// values for callers holding read without update). It sees no policies, +// guardrails, budgets, or Agent Network settings. var UsageViewer = RolePermissions{ Role: types.UserRoleUsageViewer, AutoAllowNew: map[operations.Operation]bool{ @@ -32,6 +31,12 @@ var UsageViewer = RolePermissions{ operations.Update: false, operations.Delete: false, }, + modules.AgentNetworkLogs: { + operations.Read: true, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + }, modules.AgentNetworkProviders: { operations.Read: true, operations.Create: false, diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 5a159d6f3..b705ba0a2 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -2,43 +2,23 @@ package store import ( "context" - "database/sql" - "encoding/json" - "errors" "fmt" - "math" - "net" - "net/netip" - "net/url" - "os" - "path/filepath" - "runtime" - "runtime/debug" - "strconv" - "strings" "sync" "time" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" - "github.com/rs/xid" log "github.com/sirupsen/logrus" - "gorm.io/driver/mysql" - "gorm.io/driver/postgres" - "gorm.io/driver/sqlite" "gorm.io/gorm" - "gorm.io/gorm/clause" - "gorm.io/gorm/logger" nbdns "github.com/netbirdio/netbird/dns" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" - - agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/internals/modules/zones" "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/internals/shared/db" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" @@ -46,92 +26,58 @@ import ( "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/types" - "github.com/netbirdio/netbird/management/server/util" "github.com/netbirdio/netbird/route" - "github.com/netbirdio/netbird/shared/management/status" "github.com/netbirdio/netbird/util/crypt" ) const ( - storeSqliteFileName = "store.db" idQueryCondition = "id = ?" keyQueryCondition = "key = ?" mysqlKeyQueryCondition = "`key` = ?" accountAndIDQueryCondition = "account_id = ? and id = ?" + accountAndAnyIDQueryCondition = "account_id = ? and (id = ? or public_id = ?)" accountAndPeerIDQueryCondition = "account_id = ? and peer_id = ?" accountAndIDsQueryCondition = "account_id = ? AND id IN ?" accountIDCondition = "account_id = ?" peerNotFoundFMT = "peer %s not found" - - pgMaxConnections = 30 - pgMinConnections = 1 - pgMaxConnLifetime = 60 * time.Minute - pgHealthCheckPeriod = 1 * time.Minute ) +var testPoolConfig = db.PoolConfig{ + MaxConns: 5, + MinConns: 1, + MaxConnLifetime: 30 * time.Second, + HealthCheckPeriod: 10 * time.Second, +} + // SqlStore represents an account storage backed by a Sql DB persisted to disk type SqlStore struct { - db *gorm.DB - globalAccountLock sync.Mutex - metrics telemetry.AppMetrics - installationPK int - storeEngine types.Engine - pool *pgxpool.Pool - fieldEncrypt *crypt.FieldEncrypt - transactionTimeout time.Duration -} - -type installation struct { - ID uint `gorm:"primaryKey"` - InstallationIDValue string + conn *db.Conn + db *gorm.DB + tx *db.Tx + globalAccountLock sync.Mutex + metrics telemetry.AppMetrics + installationPK int + fieldEncrypt *crypt.FieldEncrypt } type migrationFunc func(*gorm.DB) error -// NewSqlStore creates a new SqlStore instance. -func NewSqlStore(ctx context.Context, db *gorm.DB, storeEngine types.Engine, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - sql, err := db.DB() - if err != nil { - return nil, err +// NewSqlStore creates a new SqlStore instance on top of an open connection. +func NewSqlStore(ctx context.Context, conn *db.Conn, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { + if metrics != nil { + conn.SetTxMetrics(metrics.StoreMetrics()) } - - conns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS")) - if err != nil { - conns = runtime.NumCPU() - } - - transactionTimeout := 5 * time.Minute - if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" { - if parsed, err := time.ParseDuration(v); err == nil { - transactionTimeout = parsed - } - } - log.WithContext(ctx).Infof("Setting transaction timeout to %v", transactionTimeout) - - if storeEngine == types.SqliteStoreEngine { - if err == nil { - log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1") - } - conns = 1 - } - - sql.SetMaxOpenConns(conns) - sql.SetMaxIdleConns(conns) - sql.SetConnMaxLifetime(time.Hour) - sql.SetConnMaxIdleTime(3 * time.Minute) - - log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v", - conns, conns, time.Hour, 3*time.Minute) + store := &SqlStore{conn: conn, db: conn.DB(nil), metrics: metrics, installationPK: 1} if skipMigration { log.WithContext(ctx).Infof("skipping migration") - return &SqlStore{db: db, storeEngine: storeEngine, metrics: metrics, installationPK: 1, transactionTimeout: transactionTimeout}, nil + return store, nil } - if err := migratePreAuto(ctx, db); err != nil { + if err := migratePreAuto(ctx, store.db); err != nil { return nil, fmt.Errorf("migratePreAuto: %w", err) } - err = db.AutoMigrate( + err := conn.AutoMigrate( &types.SetupKey{}, &nbpeer.Peer{}, &types.User{}, &types.PersonalAccessToken{}, &types.ProxyAccessToken{}, &types.Group{}, &types.GroupPeer{}, &types.Account{}, &types.Policy{}, &types.PolicyRule{}, &route.Route{}, &nbdns.NameServerGroup{}, @@ -147,110 +93,39 @@ func NewSqlStore(ctx context.Context, db *gorm.DB, storeEngine types.Engine, met if err != nil { return nil, fmt.Errorf("auto migratePreAuto: %w", err) } - if err := migratePostAuto(ctx, db); err != nil { + if err := migratePostAuto(ctx, store.db); err != nil { return nil, fmt.Errorf("migratePostAuto: %w", err) } - return &SqlStore{db: db, storeEngine: storeEngine, metrics: metrics, installationPK: 1, transactionTimeout: transactionTimeout}, nil + return store, nil +} + +// newStore runs the migrations on conn and releases it when they fail. +func newStore(ctx context.Context, conn *db.Conn, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { + store, err := NewSqlStore(ctx, conn, metrics, skipMigration) + if err != nil { + _ = conn.Close() + return nil, err + } + return store, nil +} + +// Conn returns the shared connection so domain repositories can run alongside this store. +func (s *SqlStore) Conn() *db.Conn { + return s.conn +} + +func (s *SqlStore) pgxPool() *pgxpool.Pool { + return s.conn.Pool(s.tx) } func GetKeyQueryCondition(s *SqlStore) string { - if s.storeEngine == types.MysqlStoreEngine { + if s.conn.Engine() == db.MysqlStoreEngine { return mysqlKeyQueryCondition } return keyQueryCondition } -// SaveJob persists a job in DB -func (s *SqlStore) CreatePeerJob(ctx context.Context, job *types.Job) error { - result := s.db.Create(job) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to create job in store: %s", result.Error) - return status.Errorf(status.Internal, "failed to create job in store") - } - return nil -} - -func (s *SqlStore) CompletePeerJob(ctx context.Context, job *types.Job) error { - result := s.db. - Model(&types.Job{}). - Where(idQueryCondition, job.ID). - Updates(job) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update job in store: %s", result.Error) - return status.Errorf(status.Internal, "failed to update job in store") - } - return nil -} - -// job was pending for too long and has been cancelled -func (s *SqlStore) MarkPendingJobsAsFailed(ctx context.Context, accountID, peerID, jobID, reason string) error { - now := time.Now().UTC() - result := s.db. - Model(&types.Job{}). - Where(accountAndPeerIDQueryCondition+" AND id = ?"+" AND status = ?", accountID, peerID, jobID, types.JobStatusPending). - Updates(types.Job{ - Status: types.JobStatusFailed, - FailedReason: reason, - CompletedAt: &now, - }) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to mark pending jobs as Failed job in store: %s", result.Error) - return status.Errorf(status.Internal, "failed to mark pending job as Failed in store") - } - return nil -} - -// job was pending for too long and has been cancelled -func (s *SqlStore) MarkAllPendingJobsAsFailed(ctx context.Context, accountID, peerID, reason string) error { - now := time.Now().UTC() - result := s.db. - Model(&types.Job{}). - Where(accountAndPeerIDQueryCondition+" AND status = ?", accountID, peerID, types.JobStatusPending). - Updates(types.Job{ - Status: types.JobStatusFailed, - FailedReason: reason, - CompletedAt: &now, - }) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to mark pending jobs as Failed job in store: %s", result.Error) - return status.Errorf(status.Internal, "failed to mark pending job as Failed in store") - } - return nil -} - -// GetJobByID fetches job by ID -func (s *SqlStore) GetPeerJobByID(ctx context.Context, accountID, jobID string) (*types.Job, error) { - var job types.Job - err := s.db. - Where(accountAndIDQueryCondition, accountID, jobID). - First(&job).Error - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "job %s not found", jobID) - } - if err != nil { - log.WithContext(ctx).Errorf("failed to fetch job from store: %s", err) - return nil, err - } - return &job, nil -} - -// get all jobs -func (s *SqlStore) GetPeerJobs(ctx context.Context, accountID, peerID string) ([]*types.Job, error) { - var jobs []*types.Job - err := s.db. - Where(accountAndPeerIDQueryCondition, accountID, peerID). - Order("created_at DESC"). - Find(&jobs).Error - if err != nil { - log.WithContext(ctx).Errorf("failed to fetch jobs from store: %s", err) - return nil, err - } - - return jobs, nil -} - // AcquireGlobalLock acquires global lock across all the accounts and returns a function that releases the lock func (s *SqlStore) AcquireGlobalLock(ctx context.Context) (unlock func()) { log.WithContext(ctx).Tracef("acquiring global lock") @@ -271,2923 +146,41 @@ func (s *SqlStore) AcquireGlobalLock(ctx context.Context) (unlock func()) { return unlock } -// Deprecated: Full -// account operations are no longer supported -func (s *SqlStore) SaveAccount(ctx context.Context, account *types.Account) error { - start := time.Now() - defer func() { - elapsed := time.Since(start) - if elapsed > 1*time.Second { - log.WithContext(ctx).Tracef("SaveAccount for account %s exceeded 1s, took: %v", account.Id, elapsed) - } - }() - - // todo: remove this check after the issue is resolved - s.checkAccountDomainBeforeSave(ctx, account.Id, account.Domain) - - generateAccountSQLTypes(account) - - // Encrypt sensitive user data before saving - for i := range account.UsersG { - if err := account.UsersG[i].EncryptSensitiveData(s.fieldEncrypt); err != nil { - return fmt.Errorf("encrypt user: %w", err) - } - } - - for _, group := range account.GroupsG { - group.StoreGroupPeers() - } - - err := s.transaction(func(tx *gorm.DB) error { - result := tx.Select(clause.Associations).Delete(account.Policies, "account_id = ?", account.Id) - if result.Error != nil { - return result.Error - } - - result = tx.Select(clause.Associations).Delete(account.UsersG, "account_id = ?", account.Id) - if result.Error != nil { - return result.Error - } - - result = tx.Select(clause.Associations).Delete(account) - if result.Error != nil { - return result.Error - } - - result = tx. - Session(&gorm.Session{FullSaveAssociations: true}). - Clauses(clause.OnConflict{UpdateAll: true}). - Create(account) - if result.Error != nil { - return result.Error - } - return nil - }) - - took := time.Since(start) - if s.metrics != nil { - s.metrics.StoreMetrics().CountPersistenceDuration(took) - } - log.WithContext(ctx).Debugf("took %d ms to persist an account to the store", took.Milliseconds()) - - return err -} - -// generateAccountSQLTypes generates the GORM compatible types for the account -func generateAccountSQLTypes(account *types.Account) { - for _, key := range account.SetupKeys { - account.SetupKeysG = append(account.SetupKeysG, *key) - } - - if len(account.SetupKeys) != len(account.SetupKeysG) { - log.Warnf("SetupKeysG length mismatch for account %s", account.Id) - } - - for id, peer := range account.Peers { - peer.ID = id - account.PeersG = append(account.PeersG, *peer) - } - - for id, user := range account.Users { - user.Id = id - for id, pat := range user.PATs { - pat.ID = id - user.PATsG = append(user.PATsG, *pat) - } - account.UsersG = append(account.UsersG, *user) - } - - for id, group := range account.Groups { - group.ID = id - group.AccountID = account.Id - account.GroupsG = append(account.GroupsG, group) - } - - for id, route := range account.Routes { - route.ID = id - account.RoutesG = append(account.RoutesG, *route) - } - - for id, ns := range account.NameServerGroups { - ns.ID = id - account.NameServerGroupsG = append(account.NameServerGroupsG, *ns) - } -} - -// checkAccountDomainBeforeSave temporary method to troubleshoot an issue with domains getting blank -func (s *SqlStore) checkAccountDomainBeforeSave(ctx context.Context, accountID, newDomain string) { - var acc types.Account - var domain string - result := s.db.Model(&acc).Select("domain").Where(idQueryCondition, accountID).Take(&domain) - if result.Error != nil { - if !errors.Is(result.Error, gorm.ErrRecordNotFound) { - log.WithContext(ctx).Errorf("error when getting account %s from the store to check domain: %s", accountID, result.Error) - } - return - } - if domain != "" && newDomain == "" { - log.WithContext(ctx).Warnf("saving an account with empty domain when there was a domain set. Previous domain %s, Account ID: %s, Trace: %s", domain, accountID, debug.Stack()) - } -} - -func (s *SqlStore) DeleteAccount(ctx context.Context, account *types.Account) error { - start := time.Now() - - err := s.transaction(func(tx *gorm.DB) error { - result := tx.Select(clause.Associations).Delete(account.Policies, "account_id = ?", account.Id) - if result.Error != nil { - return result.Error - } - - result = tx.Select(clause.Associations).Delete(account.UsersG, "account_id = ?", account.Id) - if result.Error != nil { - return result.Error - } - - result = tx.Select(clause.Associations).Delete(account.Services, "account_id = ?", account.Id) - if result.Error != nil { - return result.Error - } - - result = tx.Select(clause.Associations).Delete(account) - if result.Error != nil { - return result.Error - } - - return nil - }) - - took := time.Since(start) - if s.metrics != nil { - s.metrics.StoreMetrics().CountPersistenceDuration(took) - } - log.WithContext(ctx).Tracef("took %d ms to delete an account to the store", took.Milliseconds()) - - return err -} - -func (s *SqlStore) SaveInstallationID(_ context.Context, ID string) error { - installation := installation{InstallationIDValue: ID} - installation.ID = uint(s.installationPK) - - return s.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&installation).Error -} - -func (s *SqlStore) GetInstallationID() string { - var installation installation - - if result := s.db.Take(&installation, idQueryCondition, s.installationPK); result.Error != nil { - return "" - } - - return installation.InstallationIDValue -} - -func (s *SqlStore) SavePeer(ctx context.Context, accountID string, peer *nbpeer.Peer) error { - // To maintain data integrity, we create a copy of the peer's to prevent unintended updates to other fields. - peerCopy := peer.Copy() - peerCopy.AccountID = accountID - - err := s.transaction(func(tx *gorm.DB) error { - // check if peer exists before saving - var peerID string - result := tx.Model(&nbpeer.Peer{}).Select("id").Take(&peerID, accountAndIDQueryCondition, accountID, peer.ID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return status.Errorf(status.NotFound, peerNotFoundFMT, peer.ID) - } - return result.Error - } - - if peerID == "" { - return status.Errorf(status.NotFound, peerNotFoundFMT, peer.ID) - } - - result = tx.Model(&nbpeer.Peer{}).Where(accountAndIDQueryCondition, accountID, peer.ID).Save(peerCopy) - if result.Error != nil { - return status.Errorf(status.Internal, "failed to save peer to store: %v", result.Error) - } - - return nil - }) - if err != nil { - return err - } - - return nil -} - -func (s *SqlStore) UpdateAccountDomainAttributes(ctx context.Context, accountID string, domain string, category string, isPrimaryDomain bool) error { - accountCopy := types.Account{ - Domain: domain, - DomainCategory: category, - IsDomainPrimaryAccount: isPrimaryDomain, - } - - fieldsToUpdate := []string{"domain", "domain_category", "is_domain_primary_account"} - result := s.db.Model(&types.Account{}). - Select(fieldsToUpdate). - Where(idQueryCondition, accountID). - Updates(&accountCopy) - if result.Error != nil { - return status.Errorf(status.Internal, "failed to update account domain attributes to store: %v", result.Error) - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "account %s", accountID) - } - - return nil -} - -func (s *SqlStore) SavePeerStatus(ctx context.Context, accountID, peerID string, peerStatus nbpeer.PeerStatus) error { - var peerCopy nbpeer.Peer - peerCopy.Status = &peerStatus - - fieldsToUpdate := []string{ - "peer_status_last_seen", "peer_status_session_started_at", - "peer_status_connected", "peer_status_login_expired", - "peer_status_requires_approval", - } - result := s.db.Model(&nbpeer.Peer{}). - Select(fieldsToUpdate). - Where(accountAndIDQueryCondition, accountID, peerID). - Updates(&peerCopy) - if result.Error != nil { - return status.Errorf(status.Internal, "failed to save peer status to store: %v", result.Error) - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, peerNotFoundFMT, peerID) - } - - return nil -} - -// MarkPeerConnectedIfNewerSession is an atomic optimistic-locked update. -// The peer is marked connected with the given session token only when -// the stored SessionStartedAt is strictly smaller than the incoming -// one — equivalently, when no newer stream has already taken ownership. -// The sentinel zero (set on peer creation or after a disconnect) counts -// as the smallest possible token. This is the write half of the -// fencing protocol described on PeerStatus.SessionStartedAt. -// -// The post-write side effects in the caller — geo lookup, -// schedulePeerLoginExpiration, checkAndSchedulePeerInactivityExpiration, -// OnPeersUpdated — all run AFTER this method returns and are deliberately -// outside the database write so they cannot extend the row-lock window. -// -// LastSeen is set to the database's clock (CURRENT_TIMESTAMP) at the -// moment the row is written. The caller never supplies LastSeen because -// the value would otherwise drift under lock contention — a Go-side -// time.Now() taken before the write can land minutes later than the -// actual UPDATE under load, which previously caused real ordering bugs. -func (s *SqlStore) MarkPeerConnectedIfNewerSession(ctx context.Context, accountID, peerID string, newSessionStartedAt int64) (bool, error) { - result := s.db.WithContext(ctx). - Model(&nbpeer.Peer{}). - Where(accountAndIDQueryCondition, accountID, peerID). - Where("peer_status_session_started_at < ?", newSessionStartedAt). - Updates(map[string]any{ - "peer_status_connected": true, - "peer_status_last_seen": gorm.Expr("CURRENT_TIMESTAMP"), - "peer_status_session_started_at": newSessionStartedAt, - "peer_status_login_expired": false, - }) - if result.Error != nil { - return false, status.Errorf(status.Internal, "mark peer connected: %v", result.Error) - } - return result.RowsAffected > 0, nil -} - -// MarkPeerDisconnectedIfSameSession is an atomic optimistic-locked update. -// The peer is marked disconnected only when the stored SessionStartedAt -// matches the incoming token — meaning the stream that owns the current -// session is the one ending. If a newer stream has already replaced the -// session, the update is skipped. LastSeen is set to CURRENT_TIMESTAMP at -// write time; see MarkPeerConnectedIfNewerSession for the rationale. -// -// A zero sessionStartedAt is rejected at the call site; the underlying -// WHERE on equality would otherwise match every never-connected peer. -func (s *SqlStore) MarkPeerDisconnectedIfSameSession(ctx context.Context, accountID, peerID string, sessionStartedAt int64) (bool, error) { - if sessionStartedAt == 0 { - return false, nil - } - result := s.db.WithContext(ctx). - Model(&nbpeer.Peer{}). - Where(accountAndIDQueryCondition, accountID, peerID). - Where("peer_status_session_started_at = ?", sessionStartedAt). - Updates(map[string]any{ - "peer_status_connected": false, - "peer_status_last_seen": gorm.Expr("CURRENT_TIMESTAMP"), - "peer_status_session_started_at": int64(0), - }) - if result.Error != nil { - return false, status.Errorf(status.Internal, "mark peer disconnected: %v", result.Error) - } - return result.RowsAffected > 0, nil -} - -// ApproveAccountPeers marks all peers that currently require approval in the given account as approved. -func (s *SqlStore) ApproveAccountPeers(ctx context.Context, accountID string) (int, error) { - result := s.db.Model(&nbpeer.Peer{}). - Where("account_id = ? AND peer_status_requires_approval = ?", accountID, true). - Update("peer_status_requires_approval", false) - if result.Error != nil { - return 0, status.Errorf(status.Internal, "failed to approve pending account peers: %v", result.Error) - } - - return int(result.RowsAffected), nil -} - -// RefreshPeerLastSeen updates only peer_status_last_seen. Every other status -// column is left untouched: peer_status_connected and -// peer_status_session_started_at belong to the sync stream that owns the -// session, and a blind write here would corrupt the fencing -// MarkPeerConnectedIfNewerSession relies on. -// -// LastSeen comes from the database clock for the same reason it does there: a -// Go-side timestamp is taken before the write and can land after a connect that -// used CURRENT_TIMESTAMP, dragging the column backwards. -// -// staleBefore carries the caller's throttle into the same statement, so -// concurrent requests for one peer collapse into a single write instead of -// each racing on its own stale read. The column is nullable — Status is an -// embedded pointer, so a peer stored without one leaves it NULL — and NULL -// loses every comparison, hence the explicit branch for a peer never seen. -func (s *SqlStore) RefreshPeerLastSeen(ctx context.Context, accountID, peerID string, staleBefore time.Time) (bool, error) { - result := s.db.WithContext(ctx). - Model(&nbpeer.Peer{}). - Where(accountAndIDQueryCondition, accountID, peerID). - Where("(peer_status_last_seen IS NULL OR peer_status_last_seen < ?)", staleBefore). - Update("peer_status_last_seen", gorm.Expr("CURRENT_TIMESTAMP")) - if result.Error != nil { - return false, status.Errorf(status.Internal, "refresh peer last seen: %v", result.Error) - } - - return result.RowsAffected > 0, nil -} - -// SaveUsers saves the given list of users to the database. -func (s *SqlStore) SaveUsers(ctx context.Context, users []*types.User) error { - if len(users) == 0 { - return nil - } - - usersCopy := make([]*types.User, len(users)) - for i, user := range users { - userCopy := user.Copy() - userCopy.Email = user.Email - userCopy.Name = user.Name - if err := userCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { - return fmt.Errorf("encrypt user: %w", err) - } - usersCopy[i] = userCopy - } - - result := s.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&usersCopy) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save users to store: %s", result.Error) - return status.Errorf(status.Internal, "failed to save users to store") - } - return nil -} - -// SaveUser saves the given user to the database. -func (s *SqlStore) SaveUser(ctx context.Context, user *types.User) error { - userCopy := user.Copy() - userCopy.Email = user.Email - userCopy.Name = user.Name - - if err := userCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { - return fmt.Errorf("encrypt user: %w", err) - } - - result := s.db.Save(userCopy) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save user to store: %s", result.Error) - return status.Errorf(status.Internal, "failed to save user to store") - } - return nil -} - -// CreateGroups creates the given list of groups to the database. -// groupUpsertColumns is the explicit allowlist of columns that get updated when -// CreateGroups / UpdateGroups hit a PK conflict. public_id is intentionally -// omitted so a caller passing an entity with the zero value (e.g. an HTTP -// handler-built struct) cannot reset the persisted public_id during an upsert. -// Keep this in sync with the Group schema in management/server/types/group.go. -func groupUpsertColumns() clause.Set { - return clause.AssignmentColumns([]string{ - "account_id", - "name", - "issued", - "integration_ref_id", - "integration_ref_integration_type", - "resources", - }) -} - -func (s *SqlStore) CreateGroups(ctx context.Context, accountID string, groups []*types.Group) error { - if len(groups) == 0 { - return nil - } - - return s.db.Transaction(func(tx *gorm.DB) error { - result := tx. - Clauses( - clause.OnConflict{ - Columns: []clause.Column{{Name: "id"}}, - Where: clause.Where{Exprs: []clause.Expression{clause.Eq{Column: "groups.account_id", Value: accountID}}}, - DoUpdates: groupUpsertColumns(), - }, - ). - Omit(clause.Associations). - Create(&groups) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save groups to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to save groups to store") - } - - return nil - }) -} - -// UpdateGroups updates the given list of groups to the database. -func (s *SqlStore) UpdateGroups(ctx context.Context, accountID string, groups []*types.Group) error { - if len(groups) == 0 { - return nil - } - - return s.db.Transaction(func(tx *gorm.DB) error { - result := tx. - Clauses( - clause.OnConflict{ - Columns: []clause.Column{{Name: "id"}}, - Where: clause.Where{Exprs: []clause.Expression{clause.Eq{Column: "groups.account_id", Value: accountID}}}, - DoUpdates: groupUpsertColumns(), - }, - ). - Omit(clause.Associations). - Create(&groups) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save groups to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to save groups to store") - } - - return nil - }) -} - -// DeleteHashedPAT2TokenIDIndex is noop in SqlStore -func (s *SqlStore) DeleteHashedPAT2TokenIDIndex(hashedToken string) error { - return nil -} - -// DeleteTokenID2UserIDIndex is noop in SqlStore -func (s *SqlStore) DeleteTokenID2UserIDIndex(tokenID string) error { - return nil -} - -func (s *SqlStore) GetAccountByPrivateDomain(ctx context.Context, domain string) (*types.Account, error) { - accountID, err := s.GetAccountIDByPrivateDomain(ctx, LockingStrengthNone, domain) - if err != nil { - return nil, err - } - - // TODO: rework to not call GetAccount - return s.GetAccount(ctx, accountID) -} - -func (s *SqlStore) GetAccountIDByPrivateDomain(ctx context.Context, lockStrength LockingStrength, domain string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountID string - result := tx.Model(&types.Account{}).Select("id"). - Where("domain = ? and is_domain_primary_account = ? and domain_category = ?", - strings.ToLower(domain), true, types.PrivateCategory, - ).Take(&accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.Errorf(status.NotFound, "account not found: provided domain is not registered or is not private") - } - log.WithContext(ctx).Errorf("error when getting account from the store: %s", result.Error) - return "", status.NewGetAccountFromStoreError(result.Error) - } - - return accountID, nil -} - -func (s *SqlStore) GetAccountBySetupKey(ctx context.Context, setupKey string) (*types.Account, error) { - var key types.SetupKey - result := s.db.Select("account_id").Take(&key, GetKeyQueryCondition(s), setupKey) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewSetupKeyNotFoundError(setupKey) - } - log.WithContext(ctx).Errorf("failed to get account by setup key from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get account by setup key from store") - } - - if key.AccountID == "" { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - - return s.GetAccount(ctx, key.AccountID) -} - -func (s *SqlStore) GetTokenIDByHashedToken(ctx context.Context, hashedToken string) (string, error) { - var token types.PersonalAccessToken - result := s.db.Take(&token, "hashed_token = ?", hashedToken) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.Errorf(status.NotFound, "account not found: index lookup failed") - } - log.WithContext(ctx).Errorf("error when getting token from the store: %s", result.Error) - return "", status.NewGetAccountFromStoreError(result.Error) - } - - return token.ID, nil -} - -func (s *SqlStore) GetUserByPATID(ctx context.Context, lockStrength LockingStrength, patID string) (*types.User, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var user types.User - result := tx. - Joins("JOIN personal_access_tokens ON personal_access_tokens.user_id = users.id"). - Where("personal_access_tokens.id = ?", patID).Take(&user) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewPATNotFoundError(patID) - } - log.WithContext(ctx).Errorf("failed to get token user from the store: %s", result.Error) - return nil, status.NewGetUserFromStoreError() - } - - if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt user: %w", err) - } - - return &user, nil -} - -func (s *SqlStore) GetUserByUserID(ctx context.Context, lockStrength LockingStrength, userID string) (*types.User, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var user types.User - result := tx.Take(&user, idQueryCondition, userID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewUserNotFoundError(userID) - } - return nil, status.NewGetUserFromStoreError() - } - - if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt user: %w", err) - } - - return &user, nil -} - -func (s *SqlStore) DeleteUser(ctx context.Context, accountID, userID string) error { - err := s.transaction(func(tx *gorm.DB) error { - result := tx.Delete(&types.PersonalAccessToken{}, "user_id = ?", userID) - if result.Error != nil { - return result.Error - } - - return tx.Delete(&types.User{}, accountAndIDQueryCondition, accountID, userID).Error - }) - if err != nil { - log.WithContext(ctx).Errorf("failed to delete user from the store: %s", err) - return status.Errorf(status.Internal, "failed to delete user from store") - } - - return nil -} - -func (s *SqlStore) GetAccountUsers(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.User, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var users []*types.User - result := tx.Find(&users, accountIDCondition, accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "accountID not found: index lookup failed") - } - log.WithContext(ctx).Errorf("error when getting users from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "issue getting users from store") - } - - for _, user := range users { - if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt user: %w", err) - } - } - - return users, nil -} - -func (s *SqlStore) GetAccountOwner(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.User, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var user types.User - result := tx.Take(&user, "account_id = ? AND role = ?", accountID, types.UserRoleOwner) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "account owner not found: index lookup failed") - } - return nil, status.Errorf(status.Internal, "failed to get account owner from the store") - } - - if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt user: %w", err) - } - - return &user, nil -} - -// SaveUserInvite saves a user invite to the database -func (s *SqlStore) SaveUserInvite(ctx context.Context, invite *types.UserInviteRecord) error { - inviteCopy := invite.Copy() - if err := inviteCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { - return fmt.Errorf("encrypt invite: %w", err) - } - - result := s.db.Save(inviteCopy) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save user invite to store: %s", result.Error) - return status.Errorf(status.Internal, "failed to save user invite to store") - } - return nil -} - -// GetUserInviteByID retrieves a user invite by its ID and account ID -func (s *SqlStore) GetUserInviteByID(ctx context.Context, lockStrength LockingStrength, accountID, inviteID string) (*types.UserInviteRecord, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var invite types.UserInviteRecord - result := tx.Where("account_id = ?", accountID).Take(&invite, idQueryCondition, inviteID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "user invite not found") - } - log.WithContext(ctx).Errorf("failed to get user invite from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get user invite from store") - } - - if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt invite: %w", err) - } - - return &invite, nil -} - -// GetUserInviteByHashedToken retrieves a user invite by its hashed token -func (s *SqlStore) GetUserInviteByHashedToken(ctx context.Context, lockStrength LockingStrength, hashedToken string) (*types.UserInviteRecord, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var invite types.UserInviteRecord - result := tx.Take(&invite, "hashed_token = ?", hashedToken) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "user invite not found") - } - log.WithContext(ctx).Errorf("failed to get user invite from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get user invite from store") - } - - if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt invite: %w", err) - } - - return &invite, nil -} - -// GetUserInviteByEmail retrieves a user invite by account ID and email. -// Since email is encrypted with random IVs, we fetch all invites for the account -// and compare emails in memory after decryption. -func (s *SqlStore) GetUserInviteByEmail(ctx context.Context, lockStrength LockingStrength, accountID, email string) (*types.UserInviteRecord, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var invites []*types.UserInviteRecord - result := tx.Find(&invites, "account_id = ?", accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get user invites from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get user invites from store") - } - - for _, invite := range invites { - if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt invite: %w", err) - } - if strings.EqualFold(invite.Email, email) { - return invite, nil - } - } - - return nil, status.Errorf(status.NotFound, "user invite not found for email") -} - -// GetAccountUserInvites retrieves all user invites for an account -func (s *SqlStore) GetAccountUserInvites(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.UserInviteRecord, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var invites []*types.UserInviteRecord - result := tx.Find(&invites, "account_id = ?", accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get user invites from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get user invites from store") - } - - for _, invite := range invites { - if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt invite: %w", err) - } - } - - return invites, nil -} - -// DeleteUserInvite deletes a user invite by its ID -func (s *SqlStore) DeleteUserInvite(ctx context.Context, inviteID string) error { - result := s.db.Delete(&types.UserInviteRecord{}, idQueryCondition, inviteID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete user invite from store: %s", result.Error) - return status.Errorf(status.Internal, "failed to delete user invite from store") - } - return nil -} - -func (s *SqlStore) GetAccountGroups(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.Group, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var groups []*types.Group - result := tx.Preload(clause.Associations).Find(&groups, accountIDCondition, accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "accountID not found: index lookup failed") - } - log.WithContext(ctx).Errorf("failed to get account groups from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get account groups from the store") - } - - for _, g := range groups { - g.LoadGroupPeers() - } - - return groups, nil -} - -func (s *SqlStore) GetResourceGroups(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) ([]*types.Group, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var groups []*types.Group - - likePattern := `%"ID":"` + resourceID + `"%` - - result := tx. - Preload(clause.Associations). - Where("resources LIKE ?", likePattern). - Find(&groups) - - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, nil - } - return nil, result.Error - } - - for _, g := range groups { - g.LoadGroupPeers() - } - - return groups, nil -} - -func (s *SqlStore) GetAccountsCounter(ctx context.Context) (int64, error) { - var count int64 - result := s.db.Model(&types.Account{}).Count(&count) - if result.Error != nil { - return 0, fmt.Errorf("failed to get all accounts counter: %w", result.Error) - } - - return count, nil -} - -// GetCustomDomainsCounts returns the total and validated custom domain counts. -func (s *SqlStore) GetCustomDomainsCounts(ctx context.Context) (int64, int64, error) { - var total, validated int64 - if err := s.db.Model(&domain.Domain{}).Count(&total).Error; err != nil { - return 0, 0, err - } - if err := s.db.Model(&domain.Domain{}).Where("validated = ?", true).Count(&validated).Error; err != nil { - return 0, 0, err - } - return total, validated, nil -} - -// GetProxyMetrics aggregates per-cluster + per-proxy counts for the -// self-hosted telemetry payload. Single round-trip via conditional -// aggregations so a large proxies table doesn't fan out into multiple -// queries. -func (s *SqlStore) GetProxyMetrics(ctx context.Context) (ProxyMetrics, error) { - var m ProxyMetrics - activeCutoff := time.Now().Add(-proxyActiveThreshold) - - // COUNT(DISTINCT ... CASE WHEN ...) is portable across sqlite/postgres - // (MySQL too) and keeps the round-trip to one. proxy.StatusConnected - // is the same string the cluster-capability queries use; the active - // window matches the cluster-capability semantics (only proxies - // heartbeating within ~2 * heartbeat interval count as connected). - row := s.db.WithContext(ctx). - Model(&proxy.Proxy{}). - Select( - "COUNT(DISTINCT cluster_address) AS clusters, "+ - "COUNT(DISTINCT CASE WHEN account_id IS NOT NULL THEN cluster_address END) AS clusters_byop, "+ - "COUNT(DISTINCT CASE WHEN private = ? THEN cluster_address END) AS clusters_private, "+ - "COUNT(*) AS proxies, "+ - "COUNT(CASE WHEN status = ? AND last_seen > ? THEN 1 END) AS proxies_connected", - true, - proxy.StatusConnected, - activeCutoff, - ). - Row() - if err := row.Scan(&m.Clusters, &m.ClustersBYOP, &m.ClustersPrivate, &m.Proxies, &m.ProxiesConnected); err != nil { - return ProxyMetrics{}, fmt.Errorf("scan proxy metrics: %w", err) - } - return m, nil -} - -func (s *SqlStore) GetAllAccounts(ctx context.Context) (all []*types.Account) { - var accounts []types.Account - result := s.db.Find(&accounts) - if result.Error != nil { - return all - } - - for _, account := range accounts { - if acc, err := s.GetAccount(ctx, account.Id); err == nil { - all = append(all, acc) - } - } - - return all -} - -func (s *SqlStore) GetAccountMeta(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.AccountMeta, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountMeta types.AccountMeta - result := tx.Model(&types.Account{}). - Take(&accountMeta, idQueryCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("error when getting account meta %s from the store: %s", accountID, result.Error) - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewAccountNotFoundError(accountID) - } - return nil, status.NewGetAccountFromStoreError(result.Error) - } - - return &accountMeta, nil -} - -// GetAccountOnboarding retrieves the onboarding information for a specific account. -func (s *SqlStore) GetAccountOnboarding(ctx context.Context, accountID string) (*types.AccountOnboarding, error) { - var accountOnboarding types.AccountOnboarding - result := s.db.Model(&accountOnboarding).Take(&accountOnboarding, accountIDCondition, accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewAccountOnboardingNotFoundError(accountID) - } - log.WithContext(ctx).Errorf("error when getting account onboarding %s from the store: %s", accountID, result.Error) - return nil, status.NewGetAccountFromStoreError(result.Error) - } - - return &accountOnboarding, nil -} - -// SaveAccountOnboarding updates the onboarding information for a specific account. -func (s *SqlStore) SaveAccountOnboarding(ctx context.Context, onboarding *types.AccountOnboarding) error { - result := s.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(onboarding) - if result.Error != nil { - log.WithContext(ctx).Errorf("error when saving account onboarding %s in the store: %s", onboarding.AccountID, result.Error) - return status.Errorf(status.Internal, "error when saving account onboarding %s in the store: %s", onboarding.AccountID, result.Error) - } - - return nil -} - -func (s *SqlStore) GetAccount(ctx context.Context, accountID string) (*types.Account, error) { - if s.pool != nil { - return s.getAccountPgx(ctx, accountID) - } - return s.getAccountGorm(ctx, accountID) -} - -func (s *SqlStore) getAccountGorm(ctx context.Context, accountID string) (*types.Account, error) { - start := time.Now() - defer func() { - elapsed := time.Since(start) - if elapsed > 1*time.Second { - log.WithContext(ctx).Tracef("GetAccount for account %s exceeded 1s, took: %v", accountID, elapsed) - } - }() - - var account types.Account - result := s.db.Model(&account). - Preload("UsersG.PATsG"). // have to be specified as this is nested reference - Preload("Policies.Rules"). - Preload("SetupKeysG"). - Preload("PeersG"). - Preload("UsersG"). - Preload("GroupsG.GroupPeers"). - Preload("RoutesG"). - Preload("NameServerGroupsG"). - Preload("PostureChecks"). - Preload("Networks"). - Preload("NetworkRouters"). - Preload("NetworkResources"). - Preload("Onboarding"). - Preload("Services.Targets"). - Preload("Domains"). - Take(&account, idQueryCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("error when getting account %s from the store: %s", accountID, result.Error) - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewAccountNotFoundError(accountID) - } - return nil, status.NewGetAccountFromStoreError(result.Error) - } - - account.SetupKeys = make(map[string]*types.SetupKey, len(account.SetupKeysG)) - for _, key := range account.SetupKeysG { - if key.UpdatedAt.IsZero() { - key.UpdatedAt = key.CreatedAt - } - if key.AutoGroups == nil { - key.AutoGroups = []string{} - } - account.SetupKeys[key.Key] = &key - } - account.SetupKeysG = nil - - account.Peers = make(map[string]*nbpeer.Peer, len(account.PeersG)) - for _, peer := range account.PeersG { - account.Peers[peer.ID] = &peer - } - account.PeersG = nil - account.Users = make(map[string]*types.User, len(account.UsersG)) - for _, user := range account.UsersG { - user.PATs = make(map[string]*types.PersonalAccessToken, len(user.PATs)) - for _, pat := range user.PATsG { - pat.UserID = "" - user.PATs[pat.ID] = &pat - } - if user.AutoGroups == nil { - user.AutoGroups = []string{} - } - if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt user: %w", err) - } - account.Users[user.Id] = &user - user.PATsG = nil - } - account.UsersG = nil - account.Groups = make(map[string]*types.Group, len(account.GroupsG)) - for _, group := range account.GroupsG { - group.Peers = make([]string, len(group.GroupPeers)) - for i, gp := range group.GroupPeers { - group.Peers[i] = gp.PeerID - } - if group.Resources == nil { - group.Resources = []types.Resource{} - } - account.Groups[group.ID] = group - } - account.GroupsG = nil - - account.Routes = make(map[route.ID]*route.Route, len(account.RoutesG)) - for _, route := range account.RoutesG { - account.Routes[route.ID] = &route - } - account.RoutesG = nil - account.NameServerGroups = make(map[string]*nbdns.NameServerGroup, len(account.NameServerGroupsG)) - for _, ns := range account.NameServerGroupsG { - ns.AccountID = "" - if ns.NameServers == nil { - ns.NameServers = []nbdns.NameServer{} - } - if ns.Groups == nil { - ns.Groups = []string{} - } - if ns.Domains == nil { - ns.Domains = []string{} - } - account.NameServerGroups[ns.ID] = &ns - } - account.NameServerGroupsG = nil - return &account, nil -} - -func (s *SqlStore) getAccountPgx(ctx context.Context, accountID string) (*types.Account, error) { - account, err := s.getAccount(ctx, accountID) - if err != nil { - return nil, err - } - - var wg sync.WaitGroup - errChan := make(chan error, 16) - - wg.Add(1) - go func() { - defer wg.Done() - keys, err := s.getSetupKeys(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.SetupKeysG = keys - }() - - wg.Add(1) - go func() { - defer wg.Done() - peers, err := s.getPeers(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.PeersG = peers - }() - - wg.Add(1) - go func() { - defer wg.Done() - users, err := s.getUsers(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.UsersG = users - }() - - wg.Add(1) - go func() { - defer wg.Done() - groups, err := s.getGroups(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.GroupsG = groups - }() - - wg.Add(1) - go func() { - defer wg.Done() - policies, err := s.getPolicies(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.Policies = policies - }() - - wg.Add(1) - go func() { - defer wg.Done() - routes, err := s.getRoutes(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.RoutesG = routes - }() - - wg.Add(1) - go func() { - defer wg.Done() - nsgs, err := s.getNameServerGroups(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.NameServerGroupsG = nsgs - }() - - wg.Add(1) - go func() { - defer wg.Done() - checks, err := s.getPostureChecks(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.PostureChecks = checks - }() - - wg.Add(1) - go func() { - defer wg.Done() - services, err := s.getServices(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.Services = services - }() - - wg.Add(1) - go func() { - defer wg.Done() - domains, err := s.ListCustomDomains(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.Domains = domains - }() - - wg.Add(1) - go func() { - defer wg.Done() - networks, err := s.getNetworks(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.Networks = networks - }() - - wg.Add(1) - go func() { - defer wg.Done() - routers, err := s.getNetworkRouters(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.NetworkRouters = routers - }() - - wg.Add(1) - go func() { - defer wg.Done() - resources, err := s.getNetworkResources(ctx, accountID) - if err != nil { - errChan <- err - return - } - account.NetworkResources = resources - }() - - wg.Add(1) - go func() { - defer wg.Done() - err := s.getAccountOnboarding(ctx, accountID, account) - if err != nil { - errChan <- err - return - } - }() - - wg.Wait() - close(errChan) - for e := range errChan { - if e != nil { - return nil, e - } - } - - var userIDs []string - for _, u := range account.UsersG { - userIDs = append(userIDs, u.Id) - } - var policyIDs []string - for _, p := range account.Policies { - policyIDs = append(policyIDs, p.ID) - } - var groupIDs []string - for _, g := range account.GroupsG { - groupIDs = append(groupIDs, g.ID) - } - - wg.Add(3) - errChan = make(chan error, 3) - - var pats []types.PersonalAccessToken - go func() { - defer wg.Done() - var err error - pats, err = s.getPersonalAccessTokens(ctx, userIDs) - if err != nil { - errChan <- err - } - }() - - var rules []*types.PolicyRule - go func() { - defer wg.Done() - var err error - rules, err = s.getPolicyRules(ctx, policyIDs) - if err != nil { - errChan <- err - } - }() - - var groupPeers []types.GroupPeer - go func() { - defer wg.Done() - var err error - groupPeers, err = s.getGroupPeers(ctx, groupIDs) - if err != nil { - errChan <- err - } - }() - - wg.Wait() - close(errChan) - for e := range errChan { - if e != nil { - return nil, e - } - } - - patsByUserID := make(map[string][]*types.PersonalAccessToken) - for i := range pats { - pat := &pats[i] - patsByUserID[pat.UserID] = append(patsByUserID[pat.UserID], pat) - pat.UserID = "" - } - - rulesByPolicyID := make(map[string][]*types.PolicyRule) - for _, rule := range rules { - rulesByPolicyID[rule.PolicyID] = append(rulesByPolicyID[rule.PolicyID], rule) - } - - peersByGroupID := make(map[string][]string) - for _, gp := range groupPeers { - peersByGroupID[gp.GroupID] = append(peersByGroupID[gp.GroupID], gp.PeerID) - } - - account.SetupKeys = make(map[string]*types.SetupKey, len(account.SetupKeysG)) - for i := range account.SetupKeysG { - key := &account.SetupKeysG[i] - account.SetupKeys[key.Key] = key - } - - account.Peers = make(map[string]*nbpeer.Peer, len(account.PeersG)) - for i := range account.PeersG { - peer := &account.PeersG[i] - account.Peers[peer.ID] = peer - } - - account.Users = make(map[string]*types.User, len(account.UsersG)) - for i := range account.UsersG { - user := &account.UsersG[i] - if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt user: %w", err) - } - user.PATs = make(map[string]*types.PersonalAccessToken) - if userPats, ok := patsByUserID[user.Id]; ok { - for j := range userPats { - pat := userPats[j] - user.PATs[pat.ID] = pat - } - } - account.Users[user.Id] = user - } - - for i := range account.Policies { - policy := account.Policies[i] - if policyRules, ok := rulesByPolicyID[policy.ID]; ok { - policy.Rules = policyRules - } - } - - account.Groups = make(map[string]*types.Group, len(account.GroupsG)) - for i := range account.GroupsG { - group := account.GroupsG[i] - if peerIDs, ok := peersByGroupID[group.ID]; ok { - group.Peers = peerIDs - } - account.Groups[group.ID] = group - } - - account.Routes = make(map[route.ID]*route.Route, len(account.RoutesG)) - for i := range account.RoutesG { - route := &account.RoutesG[i] - account.Routes[route.ID] = route - } - - account.NameServerGroups = make(map[string]*nbdns.NameServerGroup, len(account.NameServerGroupsG)) - for i := range account.NameServerGroupsG { - nsg := &account.NameServerGroupsG[i] - nsg.AccountID = "" - account.NameServerGroups[nsg.ID] = nsg - } - - account.SetupKeysG = nil - account.PeersG = nil - account.UsersG = nil - account.GroupsG = nil - account.RoutesG = nil - account.NameServerGroupsG = nil - - return account, nil -} - -func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Account, error) { - var account types.Account - account.Network = &types.Network{} - const accountQuery = ` - SELECT - id, created_by, created_at, domain, domain_category, is_domain_primary_account, - -- Embedded Network - network_identifier, network_net, network_net_v6, network_dns, network_serial, - -- Embedded DNSSettings - dns_settings_disabled_management_groups, - -- Embedded Settings - settings_peer_login_expiration_enabled, settings_peer_login_expiration, - settings_peer_inactivity_expiration_enabled, settings_peer_inactivity_expiration, - settings_regular_users_view_blocked, settings_groups_propagation_enabled, - settings_jwt_groups_enabled, settings_jwt_groups_claim_name, settings_jwt_allow_groups, - settings_routing_peer_dns_resolution_enabled, settings_dns_domain, settings_network_range, - settings_network_range_v6, settings_ipv6_enabled_groups, settings_lazy_connection_enabled, - settings_local_mfa_enabled, settings_metrics_push_enabled, settings_agent_network_only, - settings_dashboard_features, settings_auto_update_version, settings_auto_update_always, - settings_peer_expose_enabled, settings_peer_expose_groups, - -- Embedded ExtraSettings - settings_extra_peer_approval_enabled, settings_extra_user_approval_required, - settings_extra_integrated_validator, settings_extra_integrated_validator_groups - FROM accounts WHERE id = $1` - - var ( - sPeerLoginExpirationEnabled sql.NullBool - sPeerLoginExpiration sql.NullInt64 - sPeerInactivityExpirationEnabled sql.NullBool - sPeerInactivityExpiration sql.NullInt64 - sRegularUsersViewBlocked sql.NullBool - sGroupsPropagationEnabled sql.NullBool - sJWTGroupsEnabled sql.NullBool - sJWTGroupsClaimName sql.NullString - sJWTAllowGroups sql.NullString - sRoutingPeerDNSResolutionEnabled sql.NullBool - sDNSDomain sql.NullString - sNetworkRange sql.NullString - sNetworkRangeV6 sql.NullString - sIPv6EnabledGroups sql.NullString - sLazyConnectionEnabled sql.NullBool - sLocalMFAEnabled sql.NullBool - sMetricsPushEnabled sql.NullBool - sAgentNetworkOnly sql.NullBool - sDashboardFeatures sql.NullString - autoUpdateVersion sql.NullString - autoUpdateAlways sql.NullBool - peerExposeEnabled sql.NullBool - peerExposeGroups sql.NullString - sExtraPeerApprovalEnabled sql.NullBool - sExtraUserApprovalRequired sql.NullBool - sExtraIntegratedValidator sql.NullString - sExtraIntegratedValidatorGroups sql.NullString - networkNet sql.NullString - networkNetV6 sql.NullString - dnsSettingsDisabledGroups sql.NullString - networkIdentifier sql.NullString - networkDns sql.NullString - networkSerial sql.NullInt64 - createdAt sql.NullTime - ) - err := s.pool.QueryRow(ctx, accountQuery, accountID).Scan( - &account.Id, &account.CreatedBy, &createdAt, &account.Domain, &account.DomainCategory, &account.IsDomainPrimaryAccount, - &networkIdentifier, &networkNet, &networkNetV6, &networkDns, &networkSerial, - &dnsSettingsDisabledGroups, - &sPeerLoginExpirationEnabled, &sPeerLoginExpiration, - &sPeerInactivityExpirationEnabled, &sPeerInactivityExpiration, - &sRegularUsersViewBlocked, &sGroupsPropagationEnabled, - &sJWTGroupsEnabled, &sJWTGroupsClaimName, &sJWTAllowGroups, - &sRoutingPeerDNSResolutionEnabled, &sDNSDomain, &sNetworkRange, - &sNetworkRangeV6, &sIPv6EnabledGroups, &sLazyConnectionEnabled, - &sLocalMFAEnabled, &sMetricsPushEnabled, &sAgentNetworkOnly, - &sDashboardFeatures, &autoUpdateVersion, &autoUpdateAlways, - &peerExposeEnabled, &peerExposeGroups, - &sExtraPeerApprovalEnabled, &sExtraUserApprovalRequired, - &sExtraIntegratedValidator, &sExtraIntegratedValidatorGroups, - ) - if err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, status.NewAccountNotFoundError(accountID) - } - return nil, status.NewGetAccountFromStoreError(err) - } - - account.Settings = &types.Settings{Extra: &types.ExtraSettings{}} - if networkNet.Valid { - _ = json.Unmarshal([]byte(networkNet.String), &account.Network.Net) - } - if createdAt.Valid { - account.CreatedAt = createdAt.Time - } - if dnsSettingsDisabledGroups.Valid { - _ = json.Unmarshal([]byte(dnsSettingsDisabledGroups.String), &account.DNSSettings.DisabledManagementGroups) - } - if networkIdentifier.Valid { - account.Network.Identifier = networkIdentifier.String - } - if networkDns.Valid { - account.Network.Dns = networkDns.String - } - if networkSerial.Valid { - account.Network.Serial = uint64(networkSerial.Int64) - } - if sPeerLoginExpirationEnabled.Valid { - account.Settings.PeerLoginExpirationEnabled = sPeerLoginExpirationEnabled.Bool - } - if sPeerLoginExpiration.Valid { - account.Settings.PeerLoginExpiration = time.Duration(sPeerLoginExpiration.Int64) - } - if sPeerInactivityExpirationEnabled.Valid { - account.Settings.PeerInactivityExpirationEnabled = sPeerInactivityExpirationEnabled.Bool - } - if sPeerInactivityExpiration.Valid { - account.Settings.PeerInactivityExpiration = time.Duration(sPeerInactivityExpiration.Int64) - } - if sRegularUsersViewBlocked.Valid { - account.Settings.RegularUsersViewBlocked = sRegularUsersViewBlocked.Bool - } - if sGroupsPropagationEnabled.Valid { - account.Settings.GroupsPropagationEnabled = sGroupsPropagationEnabled.Bool - } - if sJWTGroupsEnabled.Valid { - account.Settings.JWTGroupsEnabled = sJWTGroupsEnabled.Bool - } - if sJWTGroupsClaimName.Valid { - account.Settings.JWTGroupsClaimName = sJWTGroupsClaimName.String - } - if sRoutingPeerDNSResolutionEnabled.Valid { - account.Settings.RoutingPeerDNSResolutionEnabled = sRoutingPeerDNSResolutionEnabled.Bool - } - if sDNSDomain.Valid { - account.Settings.DNSDomain = sDNSDomain.String - } - if sLazyConnectionEnabled.Valid { - account.Settings.LazyConnectionEnabled = sLazyConnectionEnabled.Bool - } - if sLocalMFAEnabled.Valid { - account.Settings.LocalMfaEnabled = sLocalMFAEnabled.Bool - } - if sMetricsPushEnabled.Valid { - account.Settings.MetricsPushEnabled = sMetricsPushEnabled.Bool - } - if sAgentNetworkOnly.Valid { - account.Settings.AgentNetworkOnly = sAgentNetworkOnly.Bool - } - if sDashboardFeatures.Valid && sDashboardFeatures.String != "" { - if err := json.Unmarshal([]byte(sDashboardFeatures.String), &account.Settings.DashboardFeatures); err != nil { - log.WithContext(ctx).Warnf("failed to unmarshal dashboard features for account %s: %v", accountID, err) - } - } - if sJWTAllowGroups.Valid { - _ = json.Unmarshal([]byte(sJWTAllowGroups.String), &account.Settings.JWTAllowGroups) - } - if sNetworkRange.Valid { - _ = json.Unmarshal([]byte(sNetworkRange.String), &account.Settings.NetworkRange) - } - if networkNetV6.Valid { - _ = json.Unmarshal([]byte(networkNetV6.String), &account.Network.NetV6) - } - if sNetworkRangeV6.Valid { - _ = json.Unmarshal([]byte(sNetworkRangeV6.String), &account.Settings.NetworkRangeV6) - } - if sIPv6EnabledGroups.Valid { - _ = json.Unmarshal([]byte(sIPv6EnabledGroups.String), &account.Settings.IPv6EnabledGroups) - } - if autoUpdateAlways.Valid { - account.Settings.AutoUpdateAlways = autoUpdateAlways.Bool - } - if autoUpdateVersion.Valid { - account.Settings.AutoUpdateVersion = autoUpdateVersion.String - } - if peerExposeEnabled.Valid { - account.Settings.PeerExposeEnabled = peerExposeEnabled.Bool - } - if peerExposeGroups.Valid { - _ = json.Unmarshal([]byte(peerExposeGroups.String), &account.Settings.PeerExposeGroups) - } - - if sExtraPeerApprovalEnabled.Valid { - account.Settings.Extra.PeerApprovalEnabled = sExtraPeerApprovalEnabled.Bool - } - if sExtraUserApprovalRequired.Valid { - account.Settings.Extra.UserApprovalRequired = sExtraUserApprovalRequired.Bool - } - if sExtraIntegratedValidator.Valid { - account.Settings.Extra.IntegratedValidator = sExtraIntegratedValidator.String - } - if sExtraIntegratedValidatorGroups.Valid { - _ = json.Unmarshal([]byte(sExtraIntegratedValidatorGroups.String), &account.Settings.Extra.IntegratedValidatorGroups) - } - return &account, nil -} - -func (s *SqlStore) getSetupKeys(ctx context.Context, accountID string) ([]types.SetupKey, error) { - const query = `SELECT id, account_id, key, key_secret, name, type, created_at, expires_at, updated_at, - revoked, used_times, last_used, auto_groups, usage_limit, ephemeral, allow_extra_dns_labels FROM setup_keys WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - - keys, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (types.SetupKey, error) { - var sk types.SetupKey - var autoGroups []byte - var skCreatedAt, expiresAt, updatedAt, lastUsed sql.NullTime - var revoked, ephemeral, allowExtraDNSLabels sql.NullBool - var usedTimes, usageLimit sql.NullInt64 - - err := row.Scan(&sk.Id, &sk.AccountID, &sk.Key, &sk.KeySecret, &sk.Name, &sk.Type, &skCreatedAt, - &expiresAt, &updatedAt, &revoked, &usedTimes, &lastUsed, &autoGroups, &usageLimit, &ephemeral, &allowExtraDNSLabels) - - if err == nil { - if expiresAt.Valid { - sk.ExpiresAt = &expiresAt.Time - } - if skCreatedAt.Valid { - sk.CreatedAt = skCreatedAt.Time - } - if updatedAt.Valid { - sk.UpdatedAt = updatedAt.Time - if sk.UpdatedAt.IsZero() { - sk.UpdatedAt = sk.CreatedAt - } - } - if lastUsed.Valid { - sk.LastUsed = &lastUsed.Time - } - if revoked.Valid { - sk.Revoked = revoked.Bool - } - if usedTimes.Valid { - sk.UsedTimes = int(usedTimes.Int64) - } - if usageLimit.Valid { - sk.UsageLimit = int(usageLimit.Int64) - } - if ephemeral.Valid { - sk.Ephemeral = ephemeral.Bool - } - if allowExtraDNSLabels.Valid { - sk.AllowExtraDNSLabels = allowExtraDNSLabels.Bool - } - if autoGroups != nil { - _ = json.Unmarshal(autoGroups, &sk.AutoGroups) - } else { - sk.AutoGroups = []string{} - } - } - return sk, err - }) - if err != nil { - return nil, err - } - return keys, nil -} - -func (s *SqlStore) getPeers(ctx context.Context, accountID string) ([]nbpeer.Peer, error) { - const query = `SELECT id, account_id, key, ip, name, dns_label, user_id, ssh_key, ssh_enabled, login_expiration_enabled, - inactivity_expiration_enabled, last_login, created_at, ephemeral, extra_dns_labels, allow_extra_dns_labels, meta_hostname, - meta_go_os, meta_kernel, meta_core, meta_platform, meta_os, meta_os_version, meta_wt_version, meta_ui_version, - meta_kernel_version, meta_network_addresses, meta_system_serial_number, meta_system_product_name, meta_system_manufacturer, - meta_environment, meta_flags, meta_files, meta_certificates, meta_capabilities, peer_status_last_seen, peer_status_session_started_at, - peer_status_connected, peer_status_login_expired, peer_status_requires_approval, location_connection_ip, - location_country_code, location_city_name, location_geo_name_id, proxy_meta_embedded, proxy_meta_cluster, ipv6, meta_sync_message_version - FROM peers WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - - peers, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nbpeer.Peer, error) { - var p nbpeer.Peer - p.Status = &nbpeer.PeerStatus{} - var ( - lastLogin, createdAt sql.NullTime - sshEnabled, loginExpirationEnabled, inactivityExpirationEnabled, ephemeral, allowExtraDNSLabels sql.NullBool - peerStatusLastSeen sql.NullTime - peerStatusSessionStartedAt sql.NullInt64 - peerStatusConnected, peerStatusLoginExpired, peerStatusRequiresApproval, proxyEmbedded sql.NullBool - ip, extraDNS, netAddr, env, flags, files, certificates, capabilities, connIP, ipv6 []byte - metaHostname, metaGoOS, metaKernel, metaCore, metaPlatform sql.NullString - metaOS, metaOSVersion, metaWtVersion, metaUIVersion, metaKernelVersion sql.NullString - metaSystemSerialNumber, metaSystemProductName, metaSystemManufacturer sql.NullString - locationCountryCode, locationCityName, proxyCluster sql.NullString - locationGeoNameID sql.NullInt64 - metaSyncMessageVersion sql.NullInt32 - ) - - err := row.Scan(&p.ID, &p.AccountID, &p.Key, &ip, &p.Name, &p.DNSLabel, &p.UserID, &p.SSHKey, &sshEnabled, - &loginExpirationEnabled, &inactivityExpirationEnabled, &lastLogin, &createdAt, &ephemeral, &extraDNS, - &allowExtraDNSLabels, &metaHostname, &metaGoOS, &metaKernel, &metaCore, &metaPlatform, - &metaOS, &metaOSVersion, &metaWtVersion, &metaUIVersion, &metaKernelVersion, &netAddr, - &metaSystemSerialNumber, &metaSystemProductName, &metaSystemManufacturer, &env, &flags, &files, &certificates, &capabilities, - &peerStatusLastSeen, &peerStatusSessionStartedAt, &peerStatusConnected, &peerStatusLoginExpired, - &peerStatusRequiresApproval, &connIP, &locationCountryCode, &locationCityName, &locationGeoNameID, - &proxyEmbedded, &proxyCluster, &ipv6, &metaSyncMessageVersion) - - if err == nil { - if lastLogin.Valid { - p.LastLogin = &lastLogin.Time - } - if createdAt.Valid { - p.CreatedAt = createdAt.Time - } - if sshEnabled.Valid { - p.SSHEnabled = sshEnabled.Bool - } - if loginExpirationEnabled.Valid { - p.LoginExpirationEnabled = loginExpirationEnabled.Bool - } - if inactivityExpirationEnabled.Valid { - p.InactivityExpirationEnabled = inactivityExpirationEnabled.Bool - } - if ephemeral.Valid { - p.Ephemeral = ephemeral.Bool - } - if allowExtraDNSLabels.Valid { - p.AllowExtraDNSLabels = allowExtraDNSLabels.Bool - } - if peerStatusLastSeen.Valid { - p.Status.LastSeen = peerStatusLastSeen.Time - } - if peerStatusSessionStartedAt.Valid { - p.Status.SessionStartedAt = peerStatusSessionStartedAt.Int64 - } - if peerStatusConnected.Valid { - p.Status.Connected = peerStatusConnected.Bool - } - if peerStatusLoginExpired.Valid { - p.Status.LoginExpired = peerStatusLoginExpired.Bool - } - if peerStatusRequiresApproval.Valid { - p.Status.RequiresApproval = peerStatusRequiresApproval.Bool - } - if metaHostname.Valid { - p.Meta.Hostname = metaHostname.String - } - if metaGoOS.Valid { - p.Meta.GoOS = metaGoOS.String - } - if metaKernel.Valid { - p.Meta.Kernel = metaKernel.String - } - if metaCore.Valid { - p.Meta.Core = metaCore.String - } - if metaPlatform.Valid { - p.Meta.Platform = metaPlatform.String - } - if metaOS.Valid { - p.Meta.OS = metaOS.String - } - if metaOSVersion.Valid { - p.Meta.OSVersion = metaOSVersion.String - } - if metaWtVersion.Valid { - p.Meta.WtVersion = metaWtVersion.String - } - if metaUIVersion.Valid { - p.Meta.UIVersion = metaUIVersion.String - } - if metaKernelVersion.Valid { - p.Meta.KernelVersion = metaKernelVersion.String - } - if metaSystemSerialNumber.Valid { - p.Meta.SystemSerialNumber = metaSystemSerialNumber.String - } - if metaSystemProductName.Valid { - p.Meta.SystemProductName = metaSystemProductName.String - } - if metaSystemManufacturer.Valid { - p.Meta.SystemManufacturer = metaSystemManufacturer.String - } - if locationCountryCode.Valid { - p.Location.CountryCode = locationCountryCode.String - } - if locationCityName.Valid { - p.Location.CityName = locationCityName.String - } - if locationGeoNameID.Valid { - p.Location.GeoNameID = uint(locationGeoNameID.Int64) - } - if proxyEmbedded.Valid { - p.ProxyMeta.Embedded = proxyEmbedded.Bool - } - if proxyCluster.Valid { - p.ProxyMeta.Cluster = proxyCluster.String - } - if ip != nil { - _ = json.Unmarshal(ip, &p.IP) - } - if ipv6 != nil { - _ = json.Unmarshal(ipv6, &p.IPv6) - } - if extraDNS != nil { - _ = json.Unmarshal(extraDNS, &p.ExtraDNSLabels) - } - if netAddr != nil { - _ = json.Unmarshal(netAddr, &p.Meta.NetworkAddresses) - } - if env != nil { - _ = json.Unmarshal(env, &p.Meta.Environment) - } - if flags != nil { - _ = json.Unmarshal(flags, &p.Meta.Flags) - } - if files != nil { - _ = json.Unmarshal(files, &p.Meta.Files) - } - if certificates != nil { - _ = json.Unmarshal(certificates, &p.Meta.Certificates) - } - if capabilities != nil { - _ = json.Unmarshal(capabilities, &p.Meta.Capabilities) - } - if connIP != nil { - _ = json.Unmarshal(connIP, &p.Location.ConnectionIP) - } - if metaSyncMessageVersion.Valid { - p.Meta.SyncMessageVersion = int(metaSyncMessageVersion.Int32) - } - } - return p, err - }) - if err != nil { - return nil, err - } - return peers, nil -} - -func (s *SqlStore) getUsers(ctx context.Context, accountID string) ([]types.User, error) { - const query = `SELECT id, account_id, role, is_service_user, non_deletable, service_user_name, auto_groups, blocked, pending_approval, last_login, created_at, issued, integration_ref_id, integration_ref_integration_type, email, name FROM users WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - users, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (types.User, error) { - var u types.User - var autoGroups []byte - var lastLogin, createdAt sql.NullTime - var isServiceUser, nonDeletable, blocked, pendingApproval sql.NullBool - err := row.Scan(&u.Id, &u.AccountID, &u.Role, &isServiceUser, &nonDeletable, &u.ServiceUserName, &autoGroups, &blocked, &pendingApproval, &lastLogin, &createdAt, &u.Issued, &u.IntegrationReference.ID, &u.IntegrationReference.IntegrationType, &u.Email, &u.Name) - if err == nil { - if lastLogin.Valid { - u.LastLogin = &lastLogin.Time - } - if createdAt.Valid { - u.CreatedAt = createdAt.Time - } - if isServiceUser.Valid { - u.IsServiceUser = isServiceUser.Bool - } - if nonDeletable.Valid { - u.NonDeletable = nonDeletable.Bool - } - if blocked.Valid { - u.Blocked = blocked.Bool - } - if pendingApproval.Valid { - u.PendingApproval = pendingApproval.Bool - } - if autoGroups != nil { - _ = json.Unmarshal(autoGroups, &u.AutoGroups) - } else { - u.AutoGroups = []string{} - } - } - return u, err - }) - if err != nil { - return nil, err - } - return users, nil -} - -func (s *SqlStore) getGroups(ctx context.Context, accountID string) ([]*types.Group, error) { - const query = `SELECT id, account_id, public_id, name, issued, resources, integration_ref_id, integration_ref_integration_type FROM groups WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - groups, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*types.Group, error) { - var g types.Group - var resources []byte - var refID sql.NullInt64 - var refType sql.NullString - err := row.Scan(&g.ID, &g.AccountID, &g.PublicID, &g.Name, &g.Issued, &resources, &refID, &refType) - if err == nil { - if refID.Valid { - g.IntegrationReference.ID = int(refID.Int64) - } - if refType.Valid { - g.IntegrationReference.IntegrationType = refType.String - } - if resources != nil { - _ = json.Unmarshal(resources, &g.Resources) - } else { - g.Resources = []types.Resource{} - } - g.GroupPeers = []types.GroupPeer{} - g.Peers = []string{} - } - return &g, err - }) - if err != nil { - return nil, err - } - return groups, nil -} - -func (s *SqlStore) getPolicies(ctx context.Context, accountID string) ([]*types.Policy, error) { - const query = `SELECT id, account_id, public_id, name, description, enabled, source_posture_checks FROM policies WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - policies, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*types.Policy, error) { - var p types.Policy - var checks []byte - var enabled sql.NullBool - err := row.Scan(&p.ID, &p.AccountID, &p.PublicID, &p.Name, &p.Description, &enabled, &checks) - if err == nil { - if enabled.Valid { - p.Enabled = enabled.Bool - } - if checks != nil { - _ = json.Unmarshal(checks, &p.SourcePostureChecks) - } - } - return &p, err - }) - if err != nil { - return nil, err - } - return policies, nil -} - -func (s *SqlStore) getRoutes(ctx context.Context, accountID string) ([]route.Route, error) { - const query = `SELECT id, account_id, public_id, network, domains, keep_route, net_id, description, peer, peer_groups, network_type, masquerade, metric, enabled, groups, access_control_groups, skip_auto_apply FROM routes WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - routes, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (route.Route, error) { - var r route.Route - var network, domains, peerGroups, groups, accessGroups []byte - var keepRoute, masquerade, enabled, skipAutoApply sql.NullBool - var metric sql.NullInt64 - err := row.Scan(&r.ID, &r.AccountID, &r.PublicID, &network, &domains, &keepRoute, &r.NetID, &r.Description, &r.Peer, &peerGroups, &r.NetworkType, &masquerade, &metric, &enabled, &groups, &accessGroups, &skipAutoApply) - if err == nil { - if keepRoute.Valid { - r.KeepRoute = keepRoute.Bool - } - if masquerade.Valid { - r.Masquerade = masquerade.Bool - } - if enabled.Valid { - r.Enabled = enabled.Bool - } - if skipAutoApply.Valid { - r.SkipAutoApply = skipAutoApply.Bool - } - if metric.Valid { - r.Metric = int(metric.Int64) - } - if network != nil { - _ = json.Unmarshal(network, &r.Network) - } - if domains != nil { - _ = json.Unmarshal(domains, &r.Domains) - } - if peerGroups != nil { - _ = json.Unmarshal(peerGroups, &r.PeerGroups) - } - if groups != nil { - _ = json.Unmarshal(groups, &r.Groups) - } - if accessGroups != nil { - _ = json.Unmarshal(accessGroups, &r.AccessControlGroups) - } - } - return r, err - }) - if err != nil { - return nil, err - } - return routes, nil -} - -func (s *SqlStore) getNameServerGroups(ctx context.Context, accountID string) ([]nbdns.NameServerGroup, error) { - const query = `SELECT id, account_id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled FROM name_server_groups WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - nsgs, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nbdns.NameServerGroup, error) { - var n nbdns.NameServerGroup - var ns, groups, domains []byte - var primary, enabled, searchDomainsEnabled sql.NullBool - err := row.Scan(&n.ID, &n.AccountID, &n.PublicID, &n.Name, &n.Description, &ns, &groups, &primary, &domains, &enabled, &searchDomainsEnabled) - if err == nil { - if primary.Valid { - n.Primary = primary.Bool - } - if enabled.Valid { - n.Enabled = enabled.Bool - } - if searchDomainsEnabled.Valid { - n.SearchDomainsEnabled = searchDomainsEnabled.Bool - } - if ns != nil { - _ = json.Unmarshal(ns, &n.NameServers) - } else { - n.NameServers = []nbdns.NameServer{} - } - if groups != nil { - _ = json.Unmarshal(groups, &n.Groups) - } else { - n.Groups = []string{} - } - if domains != nil { - _ = json.Unmarshal(domains, &n.Domains) - } else { - n.Domains = []string{} - } - } - return n, err - }) - if err != nil { - return nil, err - } - return nsgs, nil -} - -func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*posture.Checks, error) { - const query = `SELECT id, account_id, public_id, name, description, checks FROM posture_checks WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - checks, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*posture.Checks, error) { - var c posture.Checks - var checksDef []byte - err := row.Scan(&c.ID, &c.AccountID, &c.PublicID, &c.Name, &c.Description, &checksDef) - if err == nil && checksDef != nil { - _ = json.Unmarshal(checksDef, &c.Checks) - } - return &c, err - }) - if err != nil { - return nil, err - } - return checks, nil -} - -// serviceSelectColumns and targetSelectColumns are the column lists the Postgres -// pgx read path scans. They must stay in sync with the rpservice.Service and -// rpservice.Target gorm models; TestPgxServiceColumnsMatchGorm enforces this. -const serviceSelectColumns = `id, account_id, name, domain, enabled, auth, restrictions, - meta_created_at, meta_certificate_issued_at, meta_last_renewed_at, meta_status, proxy_cluster, - pass_host_header, rewrite_redirects, session_private_key, session_public_key, - mode, listen_port, port_auto_assigned, source, source_peer, terminated, - private, access_groups` - -const targetSelectColumns = `id, account_id, service_id, path, host, port, protocol, - target_id, target_type, enabled, proxy_protocol, - skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers, - direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes, - capture_content_types, agent_network, disable_access_log` - -func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { - const serviceQuery = `SELECT ` + serviceSelectColumns + ` FROM services WHERE account_id = $1` - - serviceRows, err := s.pool.Query(ctx, serviceQuery, accountID) - if err != nil { - return nil, err - } - - services, err := pgx.CollectRows(serviceRows, scanService) - if err != nil { - return nil, err - } - - if len(services) == 0 { - return services, nil - } - - serviceIDs := make([]string, len(services)) - serviceMap := make(map[string]*rpservice.Service) - for i, svc := range services { - serviceIDs[i] = svc.ID - serviceMap[svc.ID] = svc - } - - targets, err := s.getServiceTargets(ctx, serviceIDs) - if err != nil { - return nil, err - } - - for _, target := range targets { - if service, ok := serviceMap[target.ServiceID]; ok { - service.Targets = append(service.Targets, target) - } - } - - return services, nil -} - -func scanService(row pgx.CollectableRow) (*rpservice.Service, error) { - var s rpservice.Service - var auth []byte - var restrictions []byte - var accessGroups []byte - var createdAt, certIssuedAt, lastRenewedAt sql.NullTime - var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString - var mode, source, sourcePeer sql.NullString - var terminated, portAutoAssigned, private sql.NullBool - var listenPort sql.NullInt64 - err := row.Scan( - &s.ID, - &s.AccountID, - &s.Name, - &s.Domain, - &s.Enabled, - &auth, - &restrictions, - &createdAt, - &certIssuedAt, - &lastRenewedAt, - &status, - &proxyCluster, - &s.PassHostHeader, - &s.RewriteRedirects, - &sessionPrivateKey, - &sessionPublicKey, - &mode, - &listenPort, - &portAutoAssigned, - &source, - &sourcePeer, - &terminated, - &private, - &accessGroups, - ) - if err != nil { - return nil, err - } - - if auth != nil { - if err := json.Unmarshal(auth, &s.Auth); err != nil { - return nil, err - } - } - - if len(restrictions) > 0 { - if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil { - return nil, fmt.Errorf("unmarshal restrictions: %w", err) - } - } - - if len(accessGroups) > 0 { - if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { - return nil, fmt.Errorf("unmarshal access_groups: %w", err) - } - } - - if private.Valid { - s.Private = private.Bool - } - - s.Meta = serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt, status) - if proxyCluster.Valid { - s.ProxyCluster = proxyCluster.String - } - if sessionPrivateKey.Valid { - s.SessionPrivateKey = sessionPrivateKey.String - } - if sessionPublicKey.Valid { - s.SessionPublicKey = sessionPublicKey.String - } - if mode.Valid { - s.Mode = mode.String - } - if source.Valid { - s.Source = source.String - } - if sourcePeer.Valid { - s.SourcePeer = sourcePeer.String - } - if terminated.Valid { - s.Terminated = terminated.Bool - } - if portAutoAssigned.Valid { - s.PortAutoAssigned = portAutoAssigned.Bool - } - if listenPort.Valid { - if listenPort.Int64 < 0 || listenPort.Int64 > math.MaxUint16 { - return nil, fmt.Errorf("listen_port %d out of range", listenPort.Int64) - } - s.ListenPort = uint16(listenPort.Int64) - } - s.Targets = []*rpservice.Target{} - return &s, nil -} - -func serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt sql.NullTime, status sql.NullString) rpservice.Meta { - meta := rpservice.Meta{} - if createdAt.Valid { - meta.CreatedAt = createdAt.Time - } - if certIssuedAt.Valid { - t := certIssuedAt.Time - meta.CertificateIssuedAt = &t - } - if lastRenewedAt.Valid { - t := lastRenewedAt.Time - meta.LastRenewedAt = &t - } - if status.Valid { - meta.Status = status.String - } - return meta -} - -func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) { - const targetsQuery = `SELECT ` + targetSelectColumns + ` FROM targets WHERE service_id = ANY($1)` - - rows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) - if err != nil { - return nil, err - } - - return pgx.CollectRows(rows, scanTarget) -} - -func scanTarget(row pgx.CollectableRow) (*rpservice.Target, error) { - var t rpservice.Target - var path sql.NullString - var pathRewrite sql.NullString - var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool - var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64 - var customHeaders, middlewares, captureContentTypes []byte - err := row.Scan( - &t.ID, - &t.AccountID, - &t.ServiceID, - &path, - &t.Host, - &t.Port, - &t.Protocol, - &t.TargetId, - &t.TargetType, - &t.Enabled, - &proxyProtocol, - &skipTLSVerify, - &requestTimeout, - &sessionIdleTimeout, - &pathRewrite, - &customHeaders, - &directUpstream, - &middlewares, - &captureMaxRequestBytes, - &captureMaxResponseBytes, - &captureContentTypes, - &agentNetwork, - &disableAccessLog, - ) - if err != nil { - return nil, err - } - if path.Valid { - t.Path = &path.String - } - - t.ProxyProtocol = proxyProtocol.Bool - t.Options.SkipTLSVerify = skipTLSVerify.Bool - t.Options.RequestTimeout = time.Duration(requestTimeout.Int64) - t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64) - t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String) - t.Options.DirectUpstream = directUpstream.Bool - t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64 - t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64 - t.Options.AgentNetwork = agentNetwork.Bool - t.Options.DisableAccessLog = disableAccessLog.Bool - - if len(customHeaders) > 0 { - if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil { - return nil, fmt.Errorf("unmarshal custom_headers: %w", err) - } - } - if len(middlewares) > 0 { - if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil { - return nil, fmt.Errorf("unmarshal middlewares: %w", err) - } - } - if len(captureContentTypes) > 0 { - if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil { - return nil, fmt.Errorf("unmarshal capture_content_types: %w", err) - } - } - return &t, nil -} - -func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) { - const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - networks, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkTypes.Network]) - if err != nil { - return nil, err - } - result := make([]*networkTypes.Network, len(networks)) - for i := range networks { - result[i] = &networks[i] - } - return result, nil -} - -func (s *SqlStore) getNetworkRouters(ctx context.Context, accountID string) ([]*routerTypes.NetworkRouter, error) { - const query = `SELECT id, network_id, account_id, public_id, peer, peer_groups, masquerade, metric, enabled FROM network_routers WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - routers, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (routerTypes.NetworkRouter, error) { - var r routerTypes.NetworkRouter - var peerGroups []byte - var masquerade, enabled sql.NullBool - var metric sql.NullInt64 - err := row.Scan(&r.ID, &r.NetworkID, &r.AccountID, &r.PublicID, &r.Peer, &peerGroups, &masquerade, &metric, &enabled) - if err == nil { - if masquerade.Valid { - r.Masquerade = masquerade.Bool - } - if enabled.Valid { - r.Enabled = enabled.Bool - } - if metric.Valid { - r.Metric = int(metric.Int64) - } - if peerGroups != nil { - _ = json.Unmarshal(peerGroups, &r.PeerGroups) - } - } - return r, err - }) - if err != nil { - return nil, err - } - result := make([]*routerTypes.NetworkRouter, len(routers)) - for i := range routers { - result[i] = &routers[i] - } - return result, nil -} - -func (s *SqlStore) getNetworkResources(ctx context.Context, accountID string) ([]*resourceTypes.NetworkResource, error) { - const query = `SELECT id, network_id, account_id, public_id, name, description, type, domain, prefix, enabled FROM network_resources WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) - if err != nil { - return nil, err - } - resources, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (resourceTypes.NetworkResource, error) { - var r resourceTypes.NetworkResource - var prefix []byte - var enabled sql.NullBool - err := row.Scan(&r.ID, &r.NetworkID, &r.AccountID, &r.PublicID, &r.Name, &r.Description, &r.Type, &r.Domain, &prefix, &enabled) - if err == nil { - if enabled.Valid { - r.Enabled = enabled.Bool - } - if prefix != nil { - _ = json.Unmarshal(prefix, &r.Prefix) - } - } - return r, err - }) - if err != nil { - return nil, err - } - result := make([]*resourceTypes.NetworkResource, len(resources)) - for i := range resources { - result[i] = &resources[i] - } - return result, nil -} - -func (s *SqlStore) getAccountOnboarding(ctx context.Context, accountID string, account *types.Account) error { - const query = `SELECT account_id, onboarding_flow_pending, signup_form_pending, created_at, updated_at FROM account_onboardings WHERE account_id = $1` - var onboardingFlowPending, signupFormPending sql.NullBool - var createdAt, updatedAt sql.NullTime - err := s.pool.QueryRow(ctx, query, accountID).Scan( - &account.Onboarding.AccountID, - &onboardingFlowPending, - &signupFormPending, - &createdAt, - &updatedAt, - ) - if err != nil && !errors.Is(err, pgx.ErrNoRows) { - return err - } - if createdAt.Valid { - account.Onboarding.CreatedAt = createdAt.Time - } - if updatedAt.Valid { - account.Onboarding.UpdatedAt = updatedAt.Time - } - if onboardingFlowPending.Valid { - account.Onboarding.OnboardingFlowPending = onboardingFlowPending.Bool - } - if signupFormPending.Valid { - account.Onboarding.SignupFormPending = signupFormPending.Bool - } - return nil -} - -func (s *SqlStore) getPersonalAccessTokens(ctx context.Context, userIDs []string) ([]types.PersonalAccessToken, error) { - if len(userIDs) == 0 { - return nil, nil - } - const query = `SELECT id, user_id, name, hashed_token, expiration_date, created_by, created_at, last_used FROM personal_access_tokens WHERE user_id = ANY($1)` - rows, err := s.pool.Query(ctx, query, userIDs) - if err != nil { - return nil, err - } - pats, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (types.PersonalAccessToken, error) { - var pat types.PersonalAccessToken - var expirationDate, lastUsed, createdAt sql.NullTime - err := row.Scan(&pat.ID, &pat.UserID, &pat.Name, &pat.HashedToken, &expirationDate, &pat.CreatedBy, &createdAt, &lastUsed) - if err == nil { - if expirationDate.Valid { - pat.ExpirationDate = &expirationDate.Time - } - if createdAt.Valid { - pat.CreatedAt = createdAt.Time - } - if lastUsed.Valid { - pat.LastUsed = &lastUsed.Time - } - } - return pat, err - }) - if err != nil { - return nil, err - } - return pats, nil -} - -func (s *SqlStore) getPolicyRules(ctx context.Context, policyIDs []string) ([]*types.PolicyRule, error) { - if len(policyIDs) == 0 { - return nil, nil - } - const query = `SELECT id, policy_id, name, description, enabled, action, destinations, destination_resource, sources, source_resource, bidirectional, protocol, ports, port_ranges, authorized_groups, authorized_user FROM policy_rules WHERE policy_id = ANY($1)` - rows, err := s.pool.Query(ctx, query, policyIDs) - if err != nil { - return nil, err - } - rules, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*types.PolicyRule, error) { - var r types.PolicyRule - var dest, destRes, sources, sourceRes, ports, portRanges, authorizedGroups []byte - var enabled, bidirectional sql.NullBool - var authorizedUser sql.NullString - err := row.Scan(&r.ID, &r.PolicyID, &r.Name, &r.Description, &enabled, &r.Action, &dest, &destRes, &sources, &sourceRes, &bidirectional, &r.Protocol, &ports, &portRanges, &authorizedGroups, &authorizedUser) - if err == nil { - if enabled.Valid { - r.Enabled = enabled.Bool - } - if bidirectional.Valid { - r.Bidirectional = bidirectional.Bool - } - if dest != nil { - _ = json.Unmarshal(dest, &r.Destinations) - } - if destRes != nil { - _ = json.Unmarshal(destRes, &r.DestinationResource) - } - if sources != nil { - _ = json.Unmarshal(sources, &r.Sources) - } - if sourceRes != nil { - _ = json.Unmarshal(sourceRes, &r.SourceResource) - } - if ports != nil { - _ = json.Unmarshal(ports, &r.Ports) - } - if portRanges != nil { - _ = json.Unmarshal(portRanges, &r.PortRanges) - } - if authorizedGroups != nil { - _ = json.Unmarshal(authorizedGroups, &r.AuthorizedGroups) - } - if authorizedUser.Valid { - r.AuthorizedUser = authorizedUser.String - } - } - return &r, err - }) - if err != nil { - return nil, err - } - return rules, nil -} - -func (s *SqlStore) getGroupPeers(ctx context.Context, groupIDs []string) ([]types.GroupPeer, error) { - if len(groupIDs) == 0 { - return nil, nil - } - const query = `SELECT account_id, group_id, peer_id FROM group_peers WHERE group_id = ANY($1)` - rows, err := s.pool.Query(ctx, query, groupIDs) - if err != nil { - return nil, err - } - groupPeers, err := pgx.CollectRows(rows, pgx.RowToStructByName[types.GroupPeer]) - if err != nil { - return nil, err - } - return groupPeers, nil -} - -func (s *SqlStore) GetAccountByUser(ctx context.Context, userID string) (*types.Account, error) { - var user types.User - result := s.db.Select("account_id").Take(&user, idQueryCondition, userID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - return nil, status.NewGetAccountFromStoreError(result.Error) - } - - if user.AccountID == "" { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - - return s.GetAccount(ctx, user.AccountID) -} - -func (s *SqlStore) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) { - var peer nbpeer.Peer - result := s.db.Select("account_id").Take(&peer, idQueryCondition, peerID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - return nil, status.NewGetAccountFromStoreError(result.Error) - } - - if peer.AccountID == "" { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - - return s.GetAccount(ctx, peer.AccountID) -} - -func (s *SqlStore) GetAccountByPeerPubKey(ctx context.Context, peerKey string) (*types.Account, error) { - var peer nbpeer.Peer - result := s.db.Select("account_id").Take(&peer, GetKeyQueryCondition(s), peerKey) - - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - return nil, status.NewGetAccountFromStoreError(result.Error) - } - - if peer.AccountID == "" { - return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") - } - - return s.GetAccount(ctx, peer.AccountID) -} - -func (s *SqlStore) GetAnyAccountID(ctx context.Context) (string, error) { - var account types.Account - result := s.db.Select("id").Order("created_at desc").Limit(1).Find(&account) - if result.Error != nil { - return "", status.NewGetAccountFromStoreError(result.Error) - } - if result.RowsAffected == 0 { - return "", status.Errorf(status.NotFound, "account not found: index lookup failed") - } - - return account.Id, nil -} - -func (s *SqlStore) GetAccountIDByPeerPubKey(ctx context.Context, peerKey string) (string, error) { - var peer nbpeer.Peer - var accountID string - result := s.db.Model(&peer).Select("account_id").Where(GetKeyQueryCondition(s), peerKey).Take(&accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.Errorf(status.NotFound, "account not found: index lookup failed") - } - return "", status.NewGetAccountFromStoreError(result.Error) - } - - return accountID, nil -} - -func (s *SqlStore) GetAccountIDByUserID(ctx context.Context, lockStrength LockingStrength, userID string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountID string - result := tx.Model(&types.User{}). - Select("account_id").Where(idQueryCondition, userID).Take(&accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.Errorf(status.NotFound, "account not found: index lookup failed") - } - return "", status.NewGetAccountFromStoreError(result.Error) - } - - return accountID, nil -} - -func (s *SqlStore) GetAccountIDByPeerID(ctx context.Context, lockStrength LockingStrength, peerID string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountID string - result := tx.Model(&nbpeer.Peer{}). - Select("account_id").Where(idQueryCondition, peerID).Take(&accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.Errorf(status.NotFound, "peer %s account not found", peerID) - } - return "", status.NewGetAccountFromStoreError(result.Error) - } - - return accountID, nil -} - -func (s *SqlStore) GetAccountIDBySetupKey(ctx context.Context, setupKey string) (string, error) { - var accountID string - result := s.db.Model(&types.SetupKey{}).Select("account_id").Where(GetKeyQueryCondition(s), setupKey).Take(&accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.NewSetupKeyNotFoundError(setupKey) - } - log.WithContext(ctx).Errorf("failed to get account ID by setup key from store: %v", result.Error) - return "", status.Errorf(status.Internal, "failed to get account ID by setup key from store") - } - - if accountID == "" { - return "", status.Errorf(status.NotFound, "account not found: index lookup failed") - } - - return accountID, nil -} - -func (s *SqlStore) GetTakenIPs(ctx context.Context, lockStrength LockingStrength, accountID string) ([]netip.Addr, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var ipJSONStrings []string - - result := tx.Model(&nbpeer.Peer{}). - Where("account_id = ?", accountID). - Pluck("ip", &ipJSONStrings) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "no peers found for the account") - } - return nil, status.Errorf(status.Internal, "issue getting IPs from store: %s", result.Error) - } - - ips := make([]netip.Addr, len(ipJSONStrings)) - for i, ipJSON := range ipJSONStrings { - var ip netip.Addr - if err := json.Unmarshal([]byte(ipJSON), &ip); err != nil { - return nil, status.Errorf(status.Internal, "issue parsing IP JSON from store") - } - ips[i] = ip.Unmap() - } - - return ips, nil -} - -func (s *SqlStore) GetPeerLabelsInAccount(ctx context.Context, lockStrength LockingStrength, accountID string, dnsLabel string) ([]string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var labels []string - result := tx.Model(&nbpeer.Peer{}). - Where("account_id = ? AND dns_label LIKE ?", accountID, dnsLabel+"%"). - Pluck("dns_label", &labels) - - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "no peers found for the account") - } - log.WithContext(ctx).Errorf("error when getting dns labels from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "issue getting dns labels from store: %s", result.Error) - } - - return labels, nil -} - -func (s *SqlStore) GetAccountNetwork(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.Network, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountNetwork types.AccountNetwork - if err := tx.Model(&types.Account{}).Where(idQueryCondition, accountID).Take(&accountNetwork).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.NewAccountNotFoundError(accountID) - } - return nil, status.Errorf(status.Internal, "issue getting network from store: %s", err) - } - return accountNetwork.Network, nil -} - -func (s *SqlStore) GetPeerByPeerPubKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peer nbpeer.Peer - result := tx.Take(&peer, GetKeyQueryCondition(s), peerKey) - - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewPeerNotFoundError(peerKey) - } - return nil, status.Errorf(status.Internal, "issue getting peer from store: %s", result.Error) - } - - return &peer, nil -} - -func (s *SqlStore) GetAccountSettings(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.Settings, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountSettings types.AccountSettings - if err := tx.Model(&types.Account{}).Where(idQueryCondition, accountID).Take(&accountSettings).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "settings not found") - } - return nil, status.Errorf(status.Internal, "issue getting settings from store: %s", err) - } - return accountSettings.Settings, nil -} - -func (s *SqlStore) GetAccountCreatedBy(ctx context.Context, lockStrength LockingStrength, accountID string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var createdBy string - result := tx.Model(&types.Account{}). - Select("created_by").Take(&createdBy, idQueryCondition, accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.NewAccountNotFoundError(accountID) - } - return "", status.NewGetAccountFromStoreError(result.Error) - } - - return createdBy, nil -} - -// SaveUserLastLogin stores the last login time for a user in DB. -func (s *SqlStore) SaveUserLastLogin(ctx context.Context, accountID, userID string, lastLogin time.Time) error { - var user types.User - result := s.db.Take(&user, accountAndIDQueryCondition, accountID, userID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return status.NewUserNotFoundError(userID) - } - return status.NewGetUserFromStoreError() - } - - if !lastLogin.IsZero() { - user.LastLogin = &lastLogin - return s.db.Save(&user).Error - } - - return nil -} - -func (s *SqlStore) GetPostureCheckByChecksDefinition(accountID string, checks *posture.ChecksDefinition) (*posture.Checks, error) { - definitionJSON, err := json.Marshal(checks) - if err != nil { - return nil, err - } - - var postureCheck posture.Checks - err = s.db.Where("account_id = ? AND checks = ?", accountID, string(definitionJSON)).Take(&postureCheck).Error - if err != nil { - return nil, err - } - - return &postureCheck, nil -} - // Close closes the underlying DB connection func (s *SqlStore) Close(_ context.Context) error { - sql, err := s.db.DB() - if err != nil { - return fmt.Errorf("get db: %w", err) - } - return sql.Close() + return s.conn.Close() } // GetStoreEngine returns underlying store engine func (s *SqlStore) GetStoreEngine() types.Engine { - return s.storeEngine + return s.conn.Engine() } // NewSqliteStore creates a new SQLite store. func NewSqliteStore(ctx context.Context, dataDir string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - storeFile := storeSqliteFileName - if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" { - storeFile = envFile - } - - // Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc") - filePath, query, hasQuery := strings.Cut(storeFile, "?") - - connStr := filePath - if !filepath.IsAbs(filePath) { - connStr = filepath.Join(dataDir, filePath) - } - - // Compose query parameters. User-provided ?_busy_timeout (or its mattn alias - // ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at - // most that long on a lock instead of blocking the only Go-side connection. - // mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so - // the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared - // stays the default on non-Windows for the same reason as before. - parsed, _ := url.ParseQuery(query) - var defaults []string - if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" { - defaults = append(defaults, "_busy_timeout=30000") - } - if !hasQuery && runtime.GOOS != "windows" { - // To avoid `The process cannot access the file because it is being used by another process` on Windows - defaults = append(defaults, "cache=shared") - } - parts := defaults - if hasQuery { - parts = append(parts, query) - } - if len(parts) > 0 { - connStr += "?" + strings.Join(parts, "&") - } - - db, err := gorm.Open(sqlite.Open(connStr), getGormConfig()) + conn, err := db.OpenSqlite(ctx, dataDir) if err != nil { return nil, err } - - return NewSqlStore(ctx, db, types.SqliteStoreEngine, metrics, skipMigration) + return newStore(ctx, conn, metrics, skipMigration) } // NewPostgresqlStore creates a new Postgres store. func NewPostgresqlStore(ctx context.Context, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - db, err := gorm.Open(postgres.Open(dsn), getGormConfig()) + conn, err := db.OpenPostgres(ctx, dsn, db.DefaultPoolConfig) if err != nil { return nil, err } - pool, err := connectToPgDb(context.Background(), dsn) - if err != nil { - return nil, err - } - store, err := NewSqlStore(ctx, db, types.PostgresStoreEngine, metrics, skipMigration) - if err != nil { - pool.Close() - return nil, err - } - store.pool = pool - return store, nil -} - -func connectToPgDb(ctx context.Context, dsn string) (*pgxpool.Pool, error) { - config, err := pgxpool.ParseConfig(dsn) - if err != nil { - return nil, fmt.Errorf("unable to parse database config: %w", err) - } - - config.MaxConns = pgMaxConnections - config.MinConns = pgMinConnections - config.MaxConnLifetime = pgMaxConnLifetime - config.HealthCheckPeriod = pgHealthCheckPeriod - - pool, err := pgxpool.NewWithConfig(ctx, config) - if err != nil { - return nil, fmt.Errorf("unable to create connection pool: %w", err) - } - - if err := pool.Ping(ctx); err != nil { - pool.Close() - return nil, fmt.Errorf("unable to ping database: %w", err) - } - - return pool, nil + return newStore(ctx, conn, metrics, skipMigration) } // NewMysqlStore creates a new MySQL store. func NewMysqlStore(ctx context.Context, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - db, err := gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), getGormConfig()) + conn, err := db.OpenMysql(ctx, dsn) if err != nil { return nil, err } - - store, err := NewSqlStore(ctx, db, types.MysqlStoreEngine, metrics, skipMigration) - if err != nil { - closeGormDB(db) - return nil, err - } - return store, nil -} - -func getGormConfig() *gorm.Config { - return &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - CreateBatchSize: 400, - } -} - -// newPostgresStore initializes a new Postgres store. -func newPostgresStore(ctx context.Context, metrics telemetry.AppMetrics, skipMigration bool) (Store, error) { - dsn, ok := lookupDSNEnv(PostgresDsnEnv, PostgresDsnEnvLegacy) - if !ok { - return nil, fmt.Errorf("%s is not set", PostgresDsnEnv) - } - return NewPostgresqlStore(ctx, dsn, metrics, skipMigration) -} - -// newMysqlStore initializes a new MySQL store. -func newMysqlStore(ctx context.Context, metrics telemetry.AppMetrics, skipMigration bool) (Store, error) { - dsn, ok := lookupDSNEnv(mysqlDsnEnv, mysqlDsnEnvLegacy) - if !ok { - return nil, fmt.Errorf("%s is not set", mysqlDsnEnv) - } - return NewMysqlStore(ctx, dsn, metrics, skipMigration) + return newStore(ctx, conn, metrics, skipMigration) } // NewSqliteStoreFromFileStore restores a store from FileStore and stores SQLite DB in the file located in datadir. @@ -3231,7 +224,7 @@ func newPostgresqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, } if err := seedFromSqliteStore(ctx, store, sqliteStore); err != nil { - closeStore(ctx, store) + _ = store.Close(ctx) return nil, err } @@ -3240,49 +233,11 @@ func newPostgresqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, // used for tests only func NewPostgresqlStoreForTests(ctx context.Context, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - db, err := gorm.Open(postgres.Open(dsn), getGormConfig()) + conn, err := db.OpenPostgres(ctx, dsn, testPoolConfig) if err != nil { return nil, err } - pool, err := connectToPgDbForTests(context.Background(), dsn) - if err != nil { - closeGormDB(db) - return nil, err - } - store, err := NewSqlStore(ctx, db, types.PostgresStoreEngine, metrics, skipMigration) - if err != nil { - // Release the sessions, or the caller cannot drop the database. - pool.Close() - closeGormDB(db) - return nil, err - } - store.pool = pool - return store, nil -} - -// used for tests only -func connectToPgDbForTests(ctx context.Context, dsn string) (*pgxpool.Pool, error) { - config, err := pgxpool.ParseConfig(dsn) - if err != nil { - return nil, fmt.Errorf("unable to parse database config: %w", err) - } - - config.MaxConns = 5 - config.MinConns = 1 - config.MaxConnLifetime = 30 * time.Second - config.HealthCheckPeriod = 10 * time.Second - - pool, err := pgxpool.NewWithConfig(ctx, config) - if err != nil { - return nil, fmt.Errorf("unable to create connection pool: %w", err) - } - - if err := pool.Ping(ctx); err != nil { - pool.Close() - return nil, fmt.Errorf("unable to ping database: %w", err) - } - - return pool, nil + return newStore(ctx, conn, metrics, skipMigration) } // NewMysqlStoreFromSqlStore restores a store from SqlStore and stores MySQL DB. @@ -3304,15 +259,6 @@ func seedFromSqliteStore(ctx context.Context, store, sqliteStore *SqlStore) erro return nil } -// closeStore releases a store that is not handed to the caller, so a failed -// seed does not leak its connection and pool. -func closeStore(ctx context.Context, store *SqlStore) { - store.Close(ctx) - if store.pool != nil { - store.pool.Close() - } -} - func newMysqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { store, err := NewMysqlStore(ctx, dsn, metrics, skipMigration) if err != nil { @@ -3320,509 +266,41 @@ func newMysqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn s } if err := seedFromSqliteStore(ctx, store, sqliteStore); err != nil { - closeStore(ctx, store) + _ = store.Close(ctx) return nil, err } return store, nil } -func (s *SqlStore) GetSetupKeyBySecret(ctx context.Context, lockStrength LockingStrength, key string) (*types.SetupKey, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var setupKey types.SetupKey - result := tx. - Take(&setupKey, GetKeyQueryCondition(s), key) - - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.PreconditionFailed, "setup key not found") - } - log.WithContext(ctx).Errorf("failed to get setup key by secret from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get setup key by secret from store") - } - return &setupKey, nil -} - -func (s *SqlStore) IncrementSetupKeyUsage(ctx context.Context, setupKeyID string) error { - result := s.db.Model(&types.SetupKey{}). - Where(idQueryCondition, setupKeyID). - Updates(map[string]interface{}{ - "used_times": gorm.Expr("used_times + 1"), - "last_used": time.Now(), - }) - - if result.Error != nil { - return status.Errorf(status.Internal, "issue incrementing setup key usage count: %s", result.Error) - } - - if result.RowsAffected == 0 { - return status.NewSetupKeyNotFoundError(setupKeyID) - } - - return nil -} - -// AddPeerToAllGroup adds a peer to the 'All' group. Method always needs to run in a transaction -func (s *SqlStore) AddPeerToAllGroup(ctx context.Context, accountID string, peerID string) error { - var groupID string - _ = s.db.Model(types.Group{}). - Select("id"). - Where("account_id = ? AND name = ?", accountID, "All"). - Limit(1). - Scan(&groupID) - - if groupID == "" { - return status.Errorf(status.NotFound, "group 'All' not found for account %s", accountID) - } - - err := s.db.Clauses(clause.OnConflict{ - Columns: []clause.Column{{Name: "group_id"}, {Name: "peer_id"}}, - DoNothing: true, - }).Create(&types.GroupPeer{ - AccountID: accountID, - GroupID: groupID, - PeerID: peerID, - }).Error - if err != nil { - return status.Errorf(status.Internal, "error adding peer to group 'All': %v", err) - } - - return nil -} - -// AddPeerToGroup adds a peer to a group -func (s *SqlStore) AddPeerToGroup(ctx context.Context, accountID, peerID, groupID string) error { - peer := &types.GroupPeer{ - AccountID: accountID, - GroupID: groupID, - PeerID: peerID, - } - - err := s.db.Clauses(clause.OnConflict{ - Columns: []clause.Column{{Name: "group_id"}, {Name: "peer_id"}}, - DoNothing: true, - }).Create(peer).Error - if err != nil { - log.WithContext(ctx).Errorf("failed to add peer %s to group %s for account %s: %v", peerID, groupID, accountID, err) - return status.Errorf(status.Internal, "failed to add peer to group") - } - - return nil -} - -// RemovePeerFromGroup removes a peer from a group -func (s *SqlStore) RemovePeerFromGroup(ctx context.Context, peerID string, groupID string) error { - err := s.db. - Delete(&types.GroupPeer{}, "group_id = ? AND peer_id = ?", groupID, peerID).Error - if err != nil { - log.WithContext(ctx).Errorf("failed to remove peer %s from group %s: %v", peerID, groupID, err) - return status.Errorf(status.Internal, "failed to remove peer from group") - } - - return nil -} - -// RemovePeerFromAllGroups removes a peer from all groups -func (s *SqlStore) RemovePeerFromAllGroups(ctx context.Context, peerID string) error { - err := s.db. - Delete(&types.GroupPeer{}, "peer_id = ?", peerID).Error - if err != nil { - log.WithContext(ctx).Errorf("failed to remove peer %s from all groups: %v", peerID, err) - return status.Errorf(status.Internal, "failed to remove peer from all groups") - } - - return nil -} - -// AddResourceToGroup adds a resource to a group. Method always needs to run n a transaction -func (s *SqlStore) AddResourceToGroup(ctx context.Context, accountId string, groupID string, resource *types.Resource) error { - var group types.Group - result := s.db.Where(accountAndIDQueryCondition, accountId, groupID).Take(&group) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return status.NewGroupNotFoundError(groupID) - } - - return status.Errorf(status.Internal, "issue finding group: %s", result.Error) - } - - for _, res := range group.Resources { - if res.ID == resource.ID { - return nil - } - } - - group.Resources = append(group.Resources, *resource) - - if err := s.db.Save(&group).Error; err != nil { - return status.Errorf(status.Internal, "issue updating group: %s", err) - } - - return nil -} - -// RemoveResourceFromGroup removes a resource from a group. Method always needs to run in a transaction -func (s *SqlStore) RemoveResourceFromGroup(ctx context.Context, accountId string, groupID string, resourceID string) error { - var group types.Group - result := s.db.Where(accountAndIDQueryCondition, accountId, groupID).Take(&group) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return status.NewGroupNotFoundError(groupID) - } - - return status.Errorf(status.Internal, "issue finding group: %s", result.Error) - } - - for i, res := range group.Resources { - if res.ID == resourceID { - group.Resources = append(group.Resources[:i], group.Resources[i+1:]...) - break - } - } - - if err := s.db.Save(&group).Error; err != nil { - return status.Errorf(status.Internal, "issue updating group: %s", err) - } - - return nil -} - -// GetPeerGroups retrieves all groups assigned to a specific peer in a given account. -func (s *SqlStore) GetPeerGroups(ctx context.Context, lockStrength LockingStrength, accountId string, peerId string) ([]*types.Group, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var groups []*types.Group - query := tx. - Joins("JOIN group_peers ON group_peers.group_id = groups.id"). - Where("groups.account_id = ? AND group_peers.peer_id = ?", accountId, peerId). - Preload(clause.Associations). - Find(&groups) - - if query.Error != nil { - return nil, query.Error - } - - for _, group := range groups { - group.LoadGroupPeers() - } - - return groups, nil -} - -// GetPeerGroupIDs retrieves all group IDs assigned to a specific peer in a given account. -func (s *SqlStore) GetPeerGroupIDs(ctx context.Context, lockStrength LockingStrength, accountId string, peerId string) ([]string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var groupIDs []string - query := tx. - Model(&types.GroupPeer{}). - Where("account_id = ? AND peer_id = ?", accountId, peerId). - Pluck("group_id", &groupIDs) - - if query.Error != nil { - if errors.Is(query.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "no groups found for peer %s in account %s", peerId, accountId) - } - log.WithContext(ctx).Errorf("failed to get group IDs for peer %s in account %s: %v", peerId, accountId, query.Error) - return nil, status.Errorf(status.Internal, "failed to get group IDs for peer from store") - } - - return groupIDs, nil -} - -// GetAccountPeers retrieves peers for an account. -func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { - var peers []*nbpeer.Peer - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - query := tx.Where(accountIDCondition, accountID) - - if nameFilter != "" { - query = query.Where("name LIKE ?", "%"+nameFilter+"%") - } - if ipFilter != "" { - query = query.Where("ip LIKE ? OR ipv6 LIKE ?", "%"+ipFilter+"%", "%"+ipFilter+"%") - } - - if err := query.Find(&peers).Error; err != nil { - log.WithContext(ctx).Errorf("failed to get peers from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get peers from store") - } - - return peers, nil -} - -// GetUserPeers retrieves peers for a user. -func (s *SqlStore) GetUserPeers(ctx context.Context, lockStrength LockingStrength, accountID, userID string) ([]*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peers []*nbpeer.Peer - - // Exclude peers added via setup keys, as they are not user-specific and have an empty user_id. - if userID == "" { - return peers, nil - } - - result := tx. - Find(&peers, "account_id = ? AND user_id = ?", accountID, userID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get peers from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get peers from store") - } - - return peers, nil -} - -func (s *SqlStore) AddPeerToAccount(ctx context.Context, peer *nbpeer.Peer) error { - if err := s.db.Create(peer).Error; err != nil { - return status.Errorf(status.Internal, "issue adding peer to account: %s", err) - } - - return nil -} - -// GetPeerByID retrieves a peer by its ID and account ID. -func (s *SqlStore) GetPeerByID(ctx context.Context, lockStrength LockingStrength, accountID, peerID string) (*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peer *nbpeer.Peer - result := tx. - Take(&peer, accountAndIDQueryCondition, accountID, peerID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewPeerNotFoundError(peerID) - } - return nil, status.Errorf(status.Internal, "failed to get peer from store") - } - - return peer, nil -} - -// GetPeersByIDs retrieves peers by their IDs and account ID. -func (s *SqlStore) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peers []*nbpeer.Peer - result := tx.Find(&peers, accountAndIDsQueryCondition, accountID, peerIDs) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get peers by ID's from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get peers by ID's from the store") - } - - peersMap := make(map[string]*nbpeer.Peer) - for _, peer := range peers { - peersMap[peer.ID] = peer - } - - return peersMap, nil -} - -// GetAccountPeersWithExpiration retrieves a list of peers that have login expiration enabled and added by a user. -func (s *SqlStore) GetAccountPeersWithExpiration(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peers []*nbpeer.Peer - result := tx. - Where("login_expiration_enabled = ? AND peer_status_login_expired != ? AND user_id IS NOT NULL AND user_id != ''", true, true). - Find(&peers, accountIDCondition, accountID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get peers with expiration from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get peers with expiration from store") - } - - return peers, nil -} - -// GetAccountPeersWithInactivity retrieves a list of peers that have login expiration enabled and added by a user. -func (s *SqlStore) GetAccountPeersWithInactivity(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peers []*nbpeer.Peer - result := tx. - Where("inactivity_expiration_enabled = ? AND user_id IS NOT NULL AND user_id != ''", true). - Find(&peers, accountIDCondition, accountID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get peers with inactivity from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get peers with inactivity from store") - } - - return peers, nil -} - -// GetAllEphemeralPeers retrieves all peers with Ephemeral set to true across all accounts, optimized for batch processing. -func (s *SqlStore) GetAllEphemeralPeers(ctx context.Context, lockStrength LockingStrength) ([]*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var allEphemeralPeers, batchPeers []*nbpeer.Peer - result := tx. - Where("ephemeral = ?", true). - FindInBatches(&batchPeers, 1000, func(tx *gorm.DB, batch int) error { - allEphemeralPeers = append(allEphemeralPeers, batchPeers...) - return nil - }) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to retrieve ephemeral peers: %s", result.Error) - return nil, fmt.Errorf("failed to retrieve ephemeral peers") - } - - return allEphemeralPeers, nil -} - -// DeletePeer removes a peer from the store. -func (s *SqlStore) DeletePeer(ctx context.Context, accountID string, peerID string) error { - result := s.db.Delete(&nbpeer.Peer{}, accountAndIDQueryCondition, accountID, peerID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to delete peer from the store: %s", err) - return status.Errorf(status.Internal, "failed to delete peer from store") - } - - if result.RowsAffected == 0 { - return status.NewPeerNotFoundError(peerID) - } - - return nil -} - -func (s *SqlStore) IncrementNetworkSerial(ctx context.Context, accountId string) error { - result := s.db.Model(&types.Account{}).Where(idQueryCondition, accountId).Update("network_serial", gorm.Expr("network_serial + 1")) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to increment network serial count in store: %v", result.Error) - return status.Errorf(status.Internal, "failed to increment network serial count in store") - } - return nil -} - +// ExecuteInTransaction runs operation in a transaction. A store that is already +// bound to one joins it instead of opening a second, independent transaction. func (s *SqlStore) ExecuteInTransaction(ctx context.Context, operation func(store Store) error) error { - timeoutCtx, cancel := context.WithTimeout(ctx, s.transactionTimeout) - defer cancel() - - startTime := time.Now() - tx := s.db.WithContext(timeoutCtx).Begin() - if tx.Error != nil { - return tx.Error + if s.tx != nil { + return operation(s) } - defer func() { - if r := recover(); r != nil { - tx.Rollback() - panic(r) - } - }() - - if s.storeEngine == types.PostgresStoreEngine { - if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to set statement timeout: %w", err) - } - if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to set lock timeout: %w", err) - } - } - - // For MySQL, disable FK checks within this transaction to avoid deadlocks - // This is session-scoped and doesn't require SUPER privileges - if s.storeEngine == types.MysqlStoreEngine { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to disable FK checks: %w", err) - } - } - - repo := s.withTx(tx) - err := operation(repo) - if err != nil { - tx.Rollback() - if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) { - log.WithContext(ctx).Warnf("transaction exceeded %s timeout after %v, stack: %s", s.transactionTimeout, time.Since(startTime), debug.Stack()) - } - return err - } - - // Re-enable FK checks before commit (optional, as transaction end resets it) - if s.storeEngine == types.MysqlStoreEngine { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to re-enable FK checks: %w", err) - } - } - - err = tx.Commit().Error - if err != nil { - if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) { - log.WithContext(ctx).Warnf("transaction commit exceeded %s timeout after %v, stack: %s", s.transactionTimeout, time.Since(startTime), debug.Stack()) - } - return err - } - - log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime)) - if s.metrics != nil { - s.metrics.StoreMetrics().CountTransactionDuration(time.Since(startTime)) - } - - return nil + return s.conn.RunInTx(ctx, func(tx *db.Tx) error { + return operation(s.withTx(tx)) + }) } -func (s *SqlStore) withTx(tx *gorm.DB) Store { +func (s *SqlStore) withTx(tx *db.Tx) Store { return &SqlStore{ - db: tx, - storeEngine: s.storeEngine, + conn: s.conn, + db: s.conn.DB(tx), + tx: tx, fieldEncrypt: s.fieldEncrypt, } } -// transaction wraps a GORM transaction with MySQL-specific FK checks handling -// Use this instead of db.Transaction() directly to avoid deadlocks on MySQL/Aurora -func (s *SqlStore) transaction(fn func(*gorm.DB) error) error { - return s.db.Transaction(func(tx *gorm.DB) error { - // For MySQL, disable FK checks within this transaction to avoid deadlocks - // This is session-scoped and doesn't require SUPER privileges - if s.storeEngine == types.MysqlStoreEngine { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil { - return fmt.Errorf("failed to disable FK checks: %w", err) - } - } - - err := fn(tx) - - // Re-enable FK checks before commit (optional, as transaction end resets it) - if s.storeEngine == types.MysqlStoreEngine && err == nil { - if fkErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error; fkErr != nil { - return fmt.Errorf("failed to re-enable FK checks: %w", fkErr) - } - } - - return err +// transaction runs fn as a savepoint of the bound transaction, or in a new +// transaction when the store is not bound to one. +func (s *SqlStore) transaction(ctx context.Context, fn func(tx *gorm.DB) error) error { + if s.tx != nil { + return s.db.Transaction(fn) + } + return s.conn.RunInTx(ctx, func(tx *db.Tx) error { + return fn(s.conn.DB(tx)) }) } @@ -3834,2935 +312,3 @@ func (s *SqlStore) GetDB() *gorm.DB { func (s *SqlStore) SetFieldEncrypt(enc *crypt.FieldEncrypt) { s.fieldEncrypt = enc } - -func (s *SqlStore) GetAccountDNSSettings(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.DNSSettings, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountDNSSettings types.AccountDNSSettings - result := tx.Model(&types.Account{}). - Take(&accountDNSSettings, idQueryCondition, accountID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewAccountNotFoundError(accountID) - } - log.WithContext(ctx).Errorf("failed to get dns settings from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get dns settings from store") - } - return &accountDNSSettings.DNSSettings, nil -} - -// AccountExists checks whether an account exists by the given ID. -func (s *SqlStore) AccountExists(ctx context.Context, lockStrength LockingStrength, id string) (bool, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var accountID string - result := tx.Model(&types.Account{}). - Select("id").Take(&accountID, idQueryCondition, id) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return false, nil - } - return false, result.Error - } - - return accountID != "", nil -} - -// GetAccountDomainAndCategory retrieves the Domain and DomainCategory fields for an account based on the given accountID. -func (s *SqlStore) GetAccountDomainAndCategory(ctx context.Context, lockStrength LockingStrength, accountID string) (string, string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var account types.Account - result := tx.Model(&types.Account{}).Select("domain", "domain_category"). - Where(idQueryCondition, accountID).Take(&account) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", "", status.Errorf(status.NotFound, "account not found") - } - return "", "", status.Errorf(status.Internal, "failed to get domain category from store: %v", result.Error) - } - - return account.Domain, account.DomainCategory, nil -} - -// GetGroupByID retrieves a group by ID and account ID. -func (s *SqlStore) GetGroupByID(ctx context.Context, lockStrength LockingStrength, accountID, groupID string) (*types.Group, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var group *types.Group - result := tx.Preload(clause.Associations).Take(&group, accountAndIDQueryCondition, accountID, groupID) - if err := result.Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.NewGroupNotFoundError(groupID) - } - log.WithContext(ctx).Errorf("failed to get group from store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get group from store") - } - - group.LoadGroupPeers() - - return group, nil -} - -// GetGroupByName retrieves a group by name and account ID. -func (s *SqlStore) GetGroupByName(ctx context.Context, lockStrength LockingStrength, accountID, groupName string) (*types.Group, error) { - tx := s.db - - var group types.Group - - // TODO: This fix is accepted for now, but if we need to handle this more frequently - // we may need to reconsider changing the types. - query := tx.Preload(clause.Associations) - - result := query. - Model(&types.Group{}). - Joins("LEFT JOIN group_peers ON group_peers.group_id = groups.id"). - Where("groups.account_id = ? AND groups.name = ?", accountID, groupName). - Group("groups.id"). - Order("COUNT(group_peers.peer_id) DESC"). - Limit(1). - First(&group) - if err := result.Error; err != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewGroupNotFoundError(groupName) - } - log.WithContext(ctx).Errorf("failed to get group by name from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get group by name from store") - } - - group.LoadGroupPeers() - - return &group, nil -} - -// GetGroupsByIDs retrieves groups by their IDs and account ID. -func (s *SqlStore) GetGroupsByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, groupIDs []string) (map[string]*types.Group, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var groups []*types.Group - result := tx.Preload(clause.Associations).Find(&groups, accountAndIDsQueryCondition, accountID, groupIDs) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get groups by ID's from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get groups by ID's from store") - } - - groupsMap := make(map[string]*types.Group) - for _, group := range groups { - group.LoadGroupPeers() - groupsMap[group.ID] = group - } - - return groupsMap, nil -} - -// CreateGroup creates a group in the store. -func (s *SqlStore) CreateGroup(ctx context.Context, group *types.Group) error { - if group == nil { - return status.Errorf(status.InvalidArgument, "group is nil") - } - - if err := s.db.Omit(clause.Associations).Create(group).Error; err != nil { - log.WithContext(ctx).Errorf("failed to save group to store: %v", err) - return status.Errorf(status.Internal, "failed to save group to store") - } - - return nil -} - -// UpdateGroup updates a group in the store. -func (s *SqlStore) UpdateGroup(ctx context.Context, group *types.Group) error { - if group == nil { - return status.Errorf(status.InvalidArgument, "group is nil") - } - - if err := s.db.Omit(clause.Associations, "public_id").Save(group).Error; err != nil { - log.WithContext(ctx).Errorf("failed to save group to store: %v", err) - return status.Errorf(status.Internal, "failed to save group to store") - } - - return nil -} - -// DeleteGroup deletes a group from the database. -func (s *SqlStore) DeleteGroup(ctx context.Context, accountID, groupID string) error { - result := s.db.Select(clause.Associations). - Delete(&types.Group{}, accountAndIDQueryCondition, accountID, groupID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to delete group from store: %s", result.Error) - return status.Errorf(status.Internal, "failed to delete group from store") - } - - if result.RowsAffected == 0 { - return status.NewGroupNotFoundError(groupID) - } - - return nil -} - -// DeleteGroups deletes groups from the database. -func (s *SqlStore) DeleteGroups(ctx context.Context, accountID string, groupIDs []string) error { - result := s.db.Select(clause.Associations). - Delete(&types.Group{}, accountAndIDsQueryCondition, accountID, groupIDs) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete groups from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete groups from store") - } - - return nil -} - -// GetAccountPolicies retrieves policies for an account. -func (s *SqlStore) GetAccountPolicies(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.Policy, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var policies []*types.Policy - result := tx. - Preload(clause.Associations).Find(&policies, accountIDCondition, accountID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get policies from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get policies from store") - } - - return policies, nil -} - -// GetPolicyByID retrieves a policy by its ID and account ID. -func (s *SqlStore) GetPolicyByID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*types.Policy, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var policy *types.Policy - - result := tx.Preload(clause.Associations). - Take(&policy, accountAndIDQueryCondition, accountID, policyID) - if err := result.Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.NewPolicyNotFoundError(policyID) - } - log.WithContext(ctx).Errorf("failed to get policy from store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get policy from store") - } - - return policy, nil -} - -func (s *SqlStore) CreatePolicy(ctx context.Context, policy *types.Policy) error { - result := s.db.Create(policy) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to create policy in store: %s", result.Error) - return status.Errorf(status.Internal, "failed to create policy in store") - } - - return nil -} - -// SavePolicy saves a policy to the database. -func (s *SqlStore) SavePolicy(ctx context.Context, policy *types.Policy) error { - result := s.db.Session(&gorm.Session{FullSaveAssociations: true}).Omit("public_id").Save(policy) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to save policy to the store: %s", err) - return status.Errorf(status.Internal, "failed to save policy to store") - } - return nil -} - -func (s *SqlStore) DeletePolicy(ctx context.Context, accountID, policyID string) error { - return s.transaction(func(tx *gorm.DB) error { - if err := tx.Where("policy_id = ?", policyID).Delete(&types.PolicyRule{}).Error; err != nil { - return fmt.Errorf("delete policy rules: %w", err) - } - - result := tx. - Where(accountAndIDQueryCondition, accountID, policyID). - Delete(&types.Policy{}) - - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to delete policy from store: %s", err) - return status.Errorf(status.Internal, "failed to delete policy from store") - } - - if result.RowsAffected == 0 { - return status.NewPolicyNotFoundError(policyID) - } - - return nil - }) -} - -func (s *SqlStore) GetPolicyRulesByResourceID(ctx context.Context, lockStrength LockingStrength, accountID string, resourceID string) ([]*types.PolicyRule, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var policyRules []*types.PolicyRule - resourceIDPattern := `%"ID":"` + resourceID + `"%` - result := tx.Where("source_resource LIKE ? OR destination_resource LIKE ?", resourceIDPattern, resourceIDPattern). - Find(&policyRules) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get policy rules for resource id from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get policy rules for resource id from store") - } - - return policyRules, nil -} - -// GetAccountPostureChecks retrieves posture checks for an account. -func (s *SqlStore) GetAccountPostureChecks(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*posture.Checks, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var postureChecks []*posture.Checks - result := tx.Find(&postureChecks, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get posture checks from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get posture checks from store") - } - - return postureChecks, nil -} - -// GetPostureChecksByID retrieves posture checks by their ID and account ID. -func (s *SqlStore) GetPostureChecksByID(ctx context.Context, lockStrength LockingStrength, accountID, postureChecksID string) (*posture.Checks, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var postureCheck *posture.Checks - result := tx. - Take(&postureCheck, accountAndIDQueryCondition, accountID, postureChecksID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewPostureChecksNotFoundError(postureChecksID) - } - log.WithContext(ctx).Errorf("failed to get posture check from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get posture check from store") - } - - return postureCheck, nil -} - -// GetPostureChecksByIDs retrieves posture checks by their IDs and account ID. -func (s *SqlStore) GetPostureChecksByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, postureChecksIDs []string) (map[string]*posture.Checks, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var postureChecks []*posture.Checks - result := tx.Find(&postureChecks, accountAndIDsQueryCondition, accountID, postureChecksIDs) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get posture checks by ID's from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get posture checks by ID's from store") - } - - postureChecksMap := make(map[string]*posture.Checks) - for _, postureCheck := range postureChecks { - postureChecksMap[postureCheck.ID] = postureCheck - } - - return postureChecksMap, nil -} - -// SavePostureChecks saves a posture checks to the database. -func (s *SqlStore) SavePostureChecks(ctx context.Context, postureCheck *posture.Checks) error { - result := s.db.Save(postureCheck) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save posture checks to store: %s", result.Error) - return status.Errorf(status.Internal, "failed to save posture checks to store") - } - - return nil -} - -// DeletePostureChecks deletes a posture checks from the database. -func (s *SqlStore) DeletePostureChecks(ctx context.Context, accountID, postureChecksID string) error { - result := s.db.Delete(&posture.Checks{}, accountAndIDQueryCondition, accountID, postureChecksID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete posture checks from store: %s", result.Error) - return status.Errorf(status.Internal, "failed to delete posture checks from store") - } - - if result.RowsAffected == 0 { - return status.NewPostureChecksNotFoundError(postureChecksID) - } - - return nil -} - -// GetAccountRoutes retrieves network routes for an account. -func (s *SqlStore) GetAccountRoutes(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*route.Route, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var routes []*route.Route - result := tx.Find(&routes, accountIDCondition, accountID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get routes from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get routes from store") - } - - return routes, nil -} - -// GetRouteByID retrieves a route by its ID and account ID. -func (s *SqlStore) GetRouteByID(ctx context.Context, lockStrength LockingStrength, accountID string, routeID string) (*route.Route, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var route *route.Route - result := tx.Take(&route, accountAndIDQueryCondition, accountID, routeID) - if err := result.Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.NewRouteNotFoundError(routeID) - } - log.WithContext(ctx).Errorf("failed to get route from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get route from store") - } - - return route, nil -} - -// SaveRoute saves a route to the database. -func (s *SqlStore) SaveRoute(ctx context.Context, route *route.Route) error { - result := s.db.Save(route) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to save route to the store: %s", err) - return status.Errorf(status.Internal, "failed to save route to store") - } - - return nil -} - -// DeleteRoute deletes a route from the database. -func (s *SqlStore) DeleteRoute(ctx context.Context, accountID, routeID string) error { - result := s.db.Delete(&route.Route{}, accountAndIDQueryCondition, accountID, routeID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to delete route from the store: %s", err) - return status.Errorf(status.Internal, "failed to delete route from store") - } - - if result.RowsAffected == 0 { - return status.NewRouteNotFoundError(routeID) - } - - return nil -} - -// GetAccountSetupKeys retrieves setup keys for an account. -func (s *SqlStore) GetAccountSetupKeys(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.SetupKey, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var setupKeys []*types.SetupKey - result := tx. - Find(&setupKeys, accountIDCondition, accountID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get setup keys from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get setup keys from store") - } - - return setupKeys, nil -} - -// GetSetupKeyByID retrieves a setup key by its ID and account ID. -func (s *SqlStore) GetSetupKeyByID(ctx context.Context, lockStrength LockingStrength, accountID, setupKeyID string) (*types.SetupKey, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var setupKey *types.SetupKey - result := tx.Take(&setupKey, accountAndIDQueryCondition, accountID, setupKeyID) - if err := result.Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.NewSetupKeyNotFoundError(setupKeyID) - } - log.WithContext(ctx).Errorf("failed to get setup key from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get setup key from store") - } - - return setupKey, nil -} - -// SaveSetupKey saves a setup key to the database. -func (s *SqlStore) SaveSetupKey(ctx context.Context, setupKey *types.SetupKey) error { - result := s.db.Save(setupKey) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save setup key to store: %s", result.Error) - return status.Errorf(status.Internal, "failed to save setup key to store") - } - - return nil -} - -// DeleteSetupKey deletes a setup key from the database. -func (s *SqlStore) DeleteSetupKey(ctx context.Context, accountID, keyID string) error { - result := s.db.Delete(&types.SetupKey{}, accountAndIDQueryCondition, accountID, keyID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete setup key from store: %s", result.Error) - return status.Errorf(status.Internal, "failed to delete setup key from store") - } - - if result.RowsAffected == 0 { - return status.NewSetupKeyNotFoundError(keyID) - } - - return nil -} - -// GetAccountNameServerGroups retrieves name server groups for an account. -func (s *SqlStore) GetAccountNameServerGroups(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbdns.NameServerGroup, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var nsGroups []*nbdns.NameServerGroup - result := tx.Find(&nsGroups, accountIDCondition, accountID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get name server groups from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get name server groups from store") - } - - return nsGroups, nil -} - -// GetNameServerGroupByID retrieves a name server group by its ID and account ID. -func (s *SqlStore) GetNameServerGroupByID(ctx context.Context, lockStrength LockingStrength, accountID, nsGroupID string) (*nbdns.NameServerGroup, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var nsGroup *nbdns.NameServerGroup - result := tx. - Take(&nsGroup, accountAndIDQueryCondition, accountID, nsGroupID) - if err := result.Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, status.NewNameServerGroupNotFoundError(nsGroupID) - } - log.WithContext(ctx).Errorf("failed to get name server group from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get name server group from store") - } - - return nsGroup, nil -} - -// SaveNameServerGroup saves a name server group to the database. -func (s *SqlStore) SaveNameServerGroup(ctx context.Context, nameServerGroup *nbdns.NameServerGroup) error { - result := s.db.Save(nameServerGroup) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to save name server group to the store: %s", err) - return status.Errorf(status.Internal, "failed to save name server group to store") - } - return nil -} - -// DeleteNameServerGroup deletes a name server group from the database. -func (s *SqlStore) DeleteNameServerGroup(ctx context.Context, accountID, nsGroupID string) error { - result := s.db.Delete(&nbdns.NameServerGroup{}, accountAndIDQueryCondition, accountID, nsGroupID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to delete name server group from the store: %s", err) - return status.Errorf(status.Internal, "failed to delete name server group from store") - } - - if result.RowsAffected == 0 { - return status.NewNameServerGroupNotFoundError(nsGroupID) - } - - return nil -} - -// SaveDNSSettings saves the DNS settings to the store. -func (s *SqlStore) SaveDNSSettings(ctx context.Context, accountID string, settings *types.DNSSettings) error { - result := s.db.Model(&types.Account{}). - Where(idQueryCondition, accountID).Updates(&types.AccountDNSSettings{DNSSettings: *settings}) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save dns settings to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to save dns settings to store") - } - - if result.RowsAffected == 0 { - return status.NewAccountNotFoundError(accountID) - } - - return nil -} - -// SaveAccountSettings stores the account settings in DB. -func (s *SqlStore) SaveAccountSettings(ctx context.Context, accountID string, settings *types.Settings) error { - result := s.db.Model(&types.Account{}). - Select("*").Where(idQueryCondition, accountID).Updates(&types.AccountSettings{Settings: settings}) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save account settings to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to save account settings to store") - } - - // MySQL reports RowsAffected=0 for no-op updates where values don't change, - // unlike SQLite/Postgres which report matched rows. Skip the check since the - // caller (UpdateAccountSettings) already verified the account exists via - // GetAccountSettings with LockingStrengthUpdate. - - return nil -} - -func (s *SqlStore) GetAccountNetworks(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*networkTypes.Network, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var networks []*networkTypes.Network - result := tx.Find(&networks, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get networks from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get networks from store") - } - - return networks, nil -} - -func (s *SqlStore) GetNetworkByID(ctx context.Context, lockStrength LockingStrength, accountID, networkID string) (*networkTypes.Network, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var network *networkTypes.Network - result := tx.Take(&network, accountAndIDQueryCondition, accountID, networkID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewNetworkNotFoundError(networkID) - } - - log.WithContext(ctx).Errorf("failed to get network from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network from store") - } - - return network, nil -} - -func (s *SqlStore) SaveNetwork(ctx context.Context, network *networkTypes.Network) error { - result := s.db.Save(network) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save network to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to save network to store") - } - - return nil -} - -func (s *SqlStore) DeleteNetwork(ctx context.Context, accountID, networkID string) error { - result := s.db.Delete(&networkTypes.Network{}, accountAndIDQueryCondition, accountID, networkID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete network from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete network from store") - } - - if result.RowsAffected == 0 { - return status.NewNetworkNotFoundError(networkID) - } - - return nil -} - -func (s *SqlStore) GetNetworkRoutersByNetID(ctx context.Context, lockStrength LockingStrength, accountID, netID string) ([]*routerTypes.NetworkRouter, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netRouters []*routerTypes.NetworkRouter - result := tx. - Find(&netRouters, "account_id = ? AND network_id = ?", accountID, netID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get network routers from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network routers from store") - } - - return netRouters, nil -} - -func (s *SqlStore) GetNetworkRoutersByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*routerTypes.NetworkRouter, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netRouters []*routerTypes.NetworkRouter - result := tx. - Find(&netRouters, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get network routers from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network routers from store") - } - - return netRouters, nil -} - -func (s *SqlStore) GetNetworkRouterByID(ctx context.Context, lockStrength LockingStrength, accountID, routerID string) (*routerTypes.NetworkRouter, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netRouter *routerTypes.NetworkRouter - result := tx. - Take(&netRouter, accountAndIDQueryCondition, accountID, routerID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewNetworkRouterNotFoundError(routerID) - } - log.WithContext(ctx).Errorf("failed to get network router from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network router from store") - } - - return netRouter, nil -} - -func (s *SqlStore) CreateNetworkRouter(ctx context.Context, router *routerTypes.NetworkRouter) error { - if err := s.db.Create(router).Error; err != nil { - log.WithContext(ctx).Errorf("failed to create network router in store: %v", err) - return status.Errorf(status.Internal, "failed to create network router in store") - } - - return nil -} - -func (s *SqlStore) UpdateNetworkRouter(ctx context.Context, router *routerTypes.NetworkRouter) error { - result := s.db. - Select("*"). - Where(accountAndIDQueryCondition, router.AccountID, router.ID). - Updates(router) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update network router in store: %v", result.Error) - return status.Errorf(status.Internal, "failed to update network router in store") - } - - if result.RowsAffected == 0 { - return status.NewNetworkRouterNotFoundError(router.ID) - } - - return nil -} - -func (s *SqlStore) DeleteNetworkRouter(ctx context.Context, accountID, routerID string) error { - result := s.db.Delete(&routerTypes.NetworkRouter{}, accountAndIDQueryCondition, accountID, routerID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete network router from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete network router from store") - } - - if result.RowsAffected == 0 { - return status.NewNetworkRouterNotFoundError(routerID) - } - - return nil -} - -func (s *SqlStore) GetNetworkResourcesByNetID(ctx context.Context, lockStrength LockingStrength, accountID, networkID string) ([]*resourceTypes.NetworkResource, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netResources []*resourceTypes.NetworkResource - result := tx. - Find(&netResources, "account_id = ? AND network_id = ?", accountID, networkID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get network resources from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network resources from store") - } - - return netResources, nil -} - -func (s *SqlStore) GetNetworkResourcesByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*resourceTypes.NetworkResource, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netResources []*resourceTypes.NetworkResource - result := tx. - Find(&netResources, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get network resources from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network resources from store") - } - - return netResources, nil -} - -func (s *SqlStore) GetNetworkResourceByID(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) (*resourceTypes.NetworkResource, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netResources *resourceTypes.NetworkResource - result := tx. - Take(&netResources, accountAndIDQueryCondition, accountID, resourceID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewNetworkResourceNotFoundError(resourceID) - } - log.WithContext(ctx).Errorf("failed to get network resource from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network resource from store") - } - - return netResources, nil -} - -func (s *SqlStore) GetNetworkResourceByName(ctx context.Context, lockStrength LockingStrength, accountID, resourceName string) (*resourceTypes.NetworkResource, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var netResources *resourceTypes.NetworkResource - result := tx. - Take(&netResources, "account_id = ? AND name = ?", accountID, resourceName) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewNetworkResourceNotFoundError(resourceName) - } - log.WithContext(ctx).Errorf("failed to get network resource from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get network resource from store") - } - - return netResources, nil -} - -func (s *SqlStore) SaveNetworkResource(ctx context.Context, resource *resourceTypes.NetworkResource) error { - result := s.db.Save(resource) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save network resource to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to save network resource to store") - } - - return nil -} - -func (s *SqlStore) DeleteNetworkResource(ctx context.Context, accountID, resourceID string) error { - result := s.db.Delete(&resourceTypes.NetworkResource{}, accountAndIDQueryCondition, accountID, resourceID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete network resource from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete network resource from store") - } - - if result.RowsAffected == 0 { - return status.NewNetworkResourceNotFoundError(resourceID) - } - - return nil -} - -// GetPATByHashedToken returns a PersonalAccessToken by its hashed token. -func (s *SqlStore) GetPATByHashedToken(ctx context.Context, lockStrength LockingStrength, hashedToken string) (*types.PersonalAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var pat types.PersonalAccessToken - result := tx.Take(&pat, "hashed_token = ?", hashedToken) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewPATNotFoundError(hashedToken) - } - log.WithContext(ctx).Errorf("failed to get pat by hash from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get pat by hash from store") - } - - return &pat, nil -} - -// GetPATByID retrieves a personal access token by its ID and user ID. -func (s *SqlStore) GetPATByID(ctx context.Context, lockStrength LockingStrength, userID string, patID string) (*types.PersonalAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var pat types.PersonalAccessToken - result := tx. - Take(&pat, "id = ? AND user_id = ?", patID, userID) - if err := result.Error; err != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewPATNotFoundError(patID) - } - log.WithContext(ctx).Errorf("failed to get pat from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get pat from store") - } - - return &pat, nil -} - -// GetUserPATs retrieves personal access tokens for a user. -func (s *SqlStore) GetUserPATs(ctx context.Context, lockStrength LockingStrength, userID string) ([]*types.PersonalAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var pats []*types.PersonalAccessToken - result := tx.Find(&pats, "user_id = ?", userID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to get user pat's from the store: %s", err) - return nil, status.Errorf(status.Internal, "failed to get user pat's from store") - } - - return pats, nil -} - -// MarkPATUsed marks a personal access token as used. -func (s *SqlStore) MarkPATUsed(ctx context.Context, patID string) error { - patCopy := types.PersonalAccessToken{ - LastUsed: util.ToPtr(time.Now().UTC()), - } - - fieldsToUpdate := []string{"last_used"} - result := s.db.Select(fieldsToUpdate). - Where(idQueryCondition, patID).Updates(&patCopy) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to mark pat as used: %s", result.Error) - return status.Errorf(status.Internal, "failed to mark pat as used") - } - - if result.RowsAffected == 0 { - return status.NewPATNotFoundError(patID) - } - - return nil -} - -// SavePAT saves a personal access token to the database. -func (s *SqlStore) SavePAT(ctx context.Context, pat *types.PersonalAccessToken) error { - result := s.db.Save(pat) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to save pat to the store: %s", err) - return status.Errorf(status.Internal, "failed to save pat to store") - } - - return nil -} - -// DeletePAT deletes a personal access token from the database. -func (s *SqlStore) DeletePAT(ctx context.Context, userID, patID string) error { - result := s.db.Delete(&types.PersonalAccessToken{}, "user_id = ? AND id = ?", userID, patID) - if err := result.Error; err != nil { - log.WithContext(ctx).Errorf("failed to delete pat from the store: %s", err) - return status.Errorf(status.Internal, "failed to delete pat from store") - } - - if result.RowsAffected == 0 { - return status.NewPATNotFoundError(patID) - } - - return nil -} - -// GetProxyAccessTokenByHashedToken retrieves a proxy access token by its hashed value. -func (s *SqlStore) GetProxyAccessTokenByHashedToken(ctx context.Context, lockStrength LockingStrength, hashedToken types.HashedProxyToken) (*types.ProxyAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var token types.ProxyAccessToken - result := tx.Take(&token, "hashed_token = ?", hashedToken) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "proxy access token not found") - } - return nil, status.Errorf(status.Internal, "get proxy access token: %v", result.Error) - } - - return &token, nil -} - -// GetAllProxyAccessTokens retrieves all proxy access tokens. -func (s *SqlStore) GetAllProxyAccessTokens(ctx context.Context, lockStrength LockingStrength) ([]*types.ProxyAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var tokens []*types.ProxyAccessToken - result := tx.Find(&tokens) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "get proxy access tokens: %v", result.Error) - } - - return tokens, nil -} - -// SaveProxyAccessToken saves a proxy access token to the database. -func (s *SqlStore) SaveProxyAccessToken(ctx context.Context, token *types.ProxyAccessToken) error { - if result := s.db.Create(token); result.Error != nil { - return status.Errorf(status.Internal, "save proxy access token: %v", result.Error) - } - return nil -} - -// RevokeProxyAccessToken revokes a proxy access token by its ID. -func (s *SqlStore) RevokeProxyAccessToken(ctx context.Context, tokenID string) error { - result := s.db.Model(&types.ProxyAccessToken{}).Where(idQueryCondition, tokenID).Update("revoked", true) - if result.Error != nil { - return status.Errorf(status.Internal, "revoke proxy access token: %v", result.Error) - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "proxy access token not found") - } - - return nil -} - -func (s *SqlStore) GetProxyAccessTokensByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.ProxyAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var tokens []*types.ProxyAccessToken - result := tx.Where("account_id = ?", accountID).Find(&tokens) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "get proxy access tokens by account: %v", result.Error) - } - - return tokens, nil -} - -func (s *SqlStore) IsProxyAccessTokenValid(ctx context.Context, tokenID string) (bool, error) { - token, err := s.GetProxyAccessTokenByID(ctx, LockingStrengthNone, tokenID) - if err != nil { - return false, err - } - return token.IsValid(), nil -} - -func (s *SqlStore) GetProxyAccessTokenByID(ctx context.Context, lockStrength LockingStrength, tokenID string) (*types.ProxyAccessToken, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var token types.ProxyAccessToken - result := tx.Take(&token, idQueryCondition, tokenID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "proxy access token not found") - } - return nil, status.Errorf(status.Internal, "get proxy access token by ID: %v", result.Error) - } - - return &token, nil -} - -// MarkProxyAccessTokenUsed updates the last used timestamp for a proxy access token. -func (s *SqlStore) MarkProxyAccessTokenUsed(ctx context.Context, tokenID string) error { - result := s.db.Model(&types.ProxyAccessToken{}). - Where(idQueryCondition, tokenID). - Update("last_used", time.Now().UTC()) - if result.Error != nil { - return status.Errorf(status.Internal, "mark proxy access token as used: %v", result.Error) - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "proxy access token not found") - } - - return nil -} - -func (s *SqlStore) GetPeerByIP(ctx context.Context, lockStrength LockingStrength, accountID string, ip net.IP) (*nbpeer.Peer, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - column := "ip" - if ip.To4() == nil { - column = "ipv6" - } - jsonValue := fmt.Sprintf(`"%s"`, ip.String()) - - var peer nbpeer.Peer - result := tx. - Take(&peer, fmt.Sprintf("account_id = ? AND %s = ?", column), accountID, jsonValue) - if result.Error != nil { - // A tunnel-IP miss is an expected outcome (e.g. the proxy's - // ValidateTunnelPeer probing an address that isn't in the - // account roster); surface it as NotFound so callers can tell - // it apart from a real store failure. - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "peer with ip %s not found", ip.String()) - } - return nil, status.Errorf(status.Internal, "failed to get peer from store") - } - - return &peer, nil -} - -func (s *SqlStore) GetPeerIdByLabel(ctx context.Context, lockStrength LockingStrength, accountID string, hostname string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peerID string - result := tx.Model(&nbpeer.Peer{}). - Select("id"). - // Where(" = ?", hostname). - Where("account_id = ? AND dns_label = ?", accountID, hostname). - Limit(1). - Scan(&peerID) - - if peerID == "" { - return "", gorm.ErrRecordNotFound - } - - return peerID, result.Error -} - -func (s *SqlStore) CountAccountsByPrivateDomain(ctx context.Context, domain string) (int64, error) { - var count int64 - result := s.db.Model(&types.Account{}). - Where("domain = ? AND domain_category = ?", - strings.ToLower(domain), types.PrivateCategory, - ).Count(&count) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to count accounts by private domain %s: %s", domain, result.Error) - return 0, status.Errorf(status.Internal, "failed to count accounts by private domain") - } - - return count, nil -} - -func (s *SqlStore) GetAccountGroupPeers(ctx context.Context, lockStrength LockingStrength, accountID string) (map[string]map[string]struct{}, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peers []types.GroupPeer - result := tx.Find(&peers, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get account group peers from store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get account group peers from store") - } - - groupPeers := make(map[string]map[string]struct{}) - for _, peer := range peers { - if _, exists := groupPeers[peer.GroupID]; !exists { - groupPeers[peer.GroupID] = make(map[string]struct{}) - } - groupPeers[peer.GroupID][peer.PeerID] = struct{}{} - } - - return groupPeers, nil -} - -func (s *SqlStore) IsPrimaryAccount(ctx context.Context, accountID string) (bool, string, error) { - var info types.PrimaryAccountInfo - result := s.db.Model(&types.Account{}). - Select("is_domain_primary_account, domain"). - Where(idQueryCondition, accountID). - Take(&info) - - if result.Error != nil { - return false, "", status.Errorf(status.Internal, "failed to get account info: %v", result.Error) - } - - return info.IsDomainPrimaryAccount, info.Domain, nil -} - -func (s *SqlStore) MarkAccountPrimary(ctx context.Context, accountID string) error { - result := s.db.Model(&types.Account{}). - Where(idQueryCondition, accountID). - Update("is_domain_primary_account", true) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to mark account as primary: %s", result.Error) - return status.Errorf(status.Internal, "failed to mark account as primary") - } - - if result.RowsAffected == 0 { - return status.NewAccountNotFoundError(accountID) - } - - return nil -} - -type accountNetworkPatch struct { - Network *types.Network `gorm:"embedded;embeddedPrefix:network_"` -} - -func (s *SqlStore) UpdateAccountNetwork(ctx context.Context, accountID string, ipNet net.IPNet) error { - patch := accountNetworkPatch{ - Network: &types.Network{Net: ipNet}, - } - - result := s.db. - Model(&types.Account{}). - Where(idQueryCondition, accountID). - Updates(&patch) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update account network: %v", result.Error) - return status.Errorf(status.Internal, "failed to update account network") - } - if result.RowsAffected == 0 { - return status.NewAccountNotFoundError(accountID) - } - return nil -} - -// UpdateAccountNetworkV6 updates the IPv6 network range for the account. -func (s *SqlStore) UpdateAccountNetworkV6(ctx context.Context, accountID string, ipNet net.IPNet) error { - patch := accountNetworkPatch{ - Network: &types.Network{NetV6: ipNet}, - } - - result := s.db. - Model(&types.Account{}). - Where(idQueryCondition, accountID). - Updates(&patch) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update account network v6: %v", result.Error) - return status.Errorf(status.Internal, "update account network v6") - } - if result.RowsAffected == 0 { - return status.NewAccountNotFoundError(accountID) - } - return nil -} - -func (s *SqlStore) GetPeersByGroupIDs(ctx context.Context, accountID string, groupIDs []string) ([]*nbpeer.Peer, error) { - if len(groupIDs) == 0 { - return []*nbpeer.Peer{}, nil - } - - var peers []*nbpeer.Peer - peerIDsSubquery := s.db.Model(&types.GroupPeer{}). - Select("DISTINCT peer_id"). - Where("account_id = ? AND group_id IN ?", accountID, groupIDs) - - result := s.db.Where("account_id = ? AND id IN (?)", accountID, peerIDsSubquery).Find(&peers) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get peers by group IDs: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get peers by group IDs") - } - - return peers, nil -} - -func (s *SqlStore) GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error) { - if len(groupIDs) == 0 { - return nil, nil - } - - var peerIDs []string - result := s.db.Model(&types.GroupPeer{}). - Select("DISTINCT peer_id"). - Where("account_id = ? AND group_id IN ?", accountID, groupIDs). - Pluck("peer_id", &peerIDs) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "failed to get peer IDs by groups: %s", result.Error) - } - - return peerIDs, nil -} - -func (s *SqlStore) GetGroupIDsByPeerIDs(ctx context.Context, accountID string, peerIDs []string) ([]string, error) { - if len(peerIDs) == 0 { - return nil, nil - } - - var groupIDs []string - result := s.db.Model(&types.GroupPeer{}). - Select("DISTINCT group_id"). - Where("account_id = ? AND peer_id IN ?", accountID, peerIDs). - Pluck("group_id", &groupIDs) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "failed to get group IDs by peers: %s", result.Error) - } - - return groupIDs, nil -} - -// GetEmbeddedProxyPeerIDsByCluster returns peer IDs of all embedded proxy peers -// in the account, grouped by their ProxyCluster. The map is nil when no embedded -// proxy peers exist. -func (s *SqlStore) GetEmbeddedProxyPeerIDsByCluster(ctx context.Context, accountID string) (map[string][]string, error) { - type row struct { - ID string - Cluster string - } - var rows []row - result := s.db.Model(&nbpeer.Peer{}). - Select("id, proxy_meta_cluster AS cluster"). - Where("account_id = ? AND proxy_meta_embedded = ?", accountID, true). - Scan(&rows) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "failed to get embedded proxy peers: %s", result.Error) - } - - out := make(map[string][]string, len(rows)) - for _, r := range rows { - out[r.Cluster] = append(out[r.Cluster], r.ID) - } - return out, nil -} - -func (s *SqlStore) GetUserIDByPeerKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var userID string - result := tx.Model(&nbpeer.Peer{}). - Select("user_id"). - Take(&userID, GetKeyQueryCondition(s), peerKey) - - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return "", status.Errorf(status.NotFound, "peer not found: index lookup failed") - } - return "", status.Errorf(status.Internal, "failed to get user ID by peer key") - } - - return userID, nil -} - -func (s *SqlStore) CreateZone(ctx context.Context, zone *zones.Zone) error { - result := s.db.Create(zone) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to create zone to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to create zone to store") - } - - return nil -} - -func (s *SqlStore) UpdateZone(ctx context.Context, zone *zones.Zone) error { - result := s.db.Select("*").Save(zone) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update zone to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to update zone to store") - } - - return nil -} - -func (s *SqlStore) DeleteZone(ctx context.Context, accountID, zoneID string) error { - result := s.db.Delete(&zones.Zone{}, accountAndIDQueryCondition, accountID, zoneID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete zone from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete zone from store") - } - - if result.RowsAffected == 0 { - return status.NewZoneNotFoundError(zoneID) - } - - return nil -} - -func (s *SqlStore) GetZoneByID(ctx context.Context, lockStrength LockingStrength, accountID, zoneID string) (*zones.Zone, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var zone *zones.Zone - result := tx.Preload("Records").Take(&zone, accountAndIDQueryCondition, accountID, zoneID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewZoneNotFoundError(zoneID) - } - - log.WithContext(ctx).Errorf("failed to get zone from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get zone from store") - } - - return zone, nil -} - -func (s *SqlStore) GetZoneByDomain(ctx context.Context, accountID, domain string) (*zones.Zone, error) { - var zone *zones.Zone - result := s.db.Where("account_id = ? AND domain = ?", accountID, domain).First(&zone) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewZoneNotFoundError(domain) - } - - log.WithContext(ctx).Errorf("failed to get zone by domain from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get zone by domain from store") - } - - return zone, nil -} - -func (s *SqlStore) GetAccountZones(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*zones.Zone, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var zones []*zones.Zone - result := tx.Preload("Records").Find(&zones, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get zones from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get zones from store") - } - - return zones, nil -} - -func (s *SqlStore) CreateDNSRecord(ctx context.Context, record *records.Record) error { - result := s.db.Create(record) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to create dns record to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to create dns record to store") - } - - return nil -} - -func (s *SqlStore) UpdateDNSRecord(ctx context.Context, record *records.Record) error { - result := s.db.Select("*").Save(record) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update dns record to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to update dns record to store") - } - - return nil -} - -func (s *SqlStore) DeleteDNSRecord(ctx context.Context, accountID, zoneID, recordID string) error { - result := s.db.Delete(&records.Record{}, "account_id = ? AND zone_id = ? AND id = ?", accountID, zoneID, recordID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete dns record from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete dns record from store") - } - - if result.RowsAffected == 0 { - return status.NewDNSRecordNotFoundError(recordID) - } - - return nil -} - -func (s *SqlStore) GetDNSRecordByID(ctx context.Context, lockStrength LockingStrength, accountID, zoneID, recordID string) (*records.Record, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var record *records.Record - result := tx.Where("account_id = ? AND zone_id = ? AND id = ?", accountID, zoneID, recordID).Take(&record) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.NewDNSRecordNotFoundError(recordID) - } - - log.WithContext(ctx).Errorf("failed to get dns record from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get dns record from store") - } - - return record, nil -} - -func (s *SqlStore) GetZoneDNSRecords(ctx context.Context, lockStrength LockingStrength, accountID, zoneID string) ([]*records.Record, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var recordsList []*records.Record - result := tx.Where("account_id = ? AND zone_id = ?", accountID, zoneID).Find(&recordsList) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get zone dns records from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get zone dns records from store") - } - - return recordsList, nil -} - -func (s *SqlStore) GetZoneDNSRecordsByName(ctx context.Context, lockStrength LockingStrength, accountID, zoneID, name string) ([]*records.Record, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var recordsList []*records.Record - result := tx.Where("account_id = ? AND zone_id = ? AND name = ?", accountID, zoneID, name).Find(&recordsList) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get zone dns records by name from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get zone dns records by name from store") - } - - return recordsList, nil -} - -func (s *SqlStore) DeleteZoneDNSRecords(ctx context.Context, accountID, zoneID string) error { - result := s.db.Delete(&records.Record{}, "account_id = ? AND zone_id = ?", accountID, zoneID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete zone dns records from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete zone dns records from store") - } - - return nil -} - -func (s *SqlStore) GetPeerIDByKey(ctx context.Context, lockStrength LockingStrength, key string) (string, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var peerID string - result := tx.Model(&nbpeer.Peer{}). - Select("id"). - Where(GetKeyQueryCondition(s), key). - Limit(1). - Scan(&peerID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get peer ID by key: %s", result.Error) - return "", status.Errorf(status.Internal, "failed to get peer ID by key") - } - - return peerID, nil -} - -func (s *SqlStore) CreateService(ctx context.Context, service *rpservice.Service) error { - serviceCopy := service.Copy() - if err := serviceCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { - return fmt.Errorf("encrypt service data: %w", err) - } - result := s.db.Create(serviceCopy) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to create service to store: %v", result.Error) - return status.Errorf(status.Internal, "failed to create service to store") - } - - return nil -} - -func (s *SqlStore) UpdateService(ctx context.Context, service *rpservice.Service) error { - serviceCopy := service.Copy() - if err := serviceCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { - return fmt.Errorf("encrypt service data: %w", err) - } - - // Create target type instance outside transaction to avoid variable shadowing - targetType := &rpservice.Target{} - - // Use a transaction to ensure atomic updates of the service and its targets - err := s.db.Transaction(func(tx *gorm.DB) error { - // Delete existing targets - if err := tx.Where("service_id = ?", serviceCopy.ID).Delete(targetType).Error; err != nil { - return err - } - - // Update the service and create new targets - if err := tx.Session(&gorm.Session{FullSaveAssociations: true}).Save(serviceCopy).Error; err != nil { - return err - } - - return nil - }) - if err != nil { - log.WithContext(ctx).Errorf("failed to update service to store: %v", err) - return status.Errorf(status.Internal, "failed to update service to store") - } - - return nil -} - -func (s *SqlStore) DeleteService(ctx context.Context, accountID, serviceID string) error { - result := s.db.Delete(&rpservice.Service{}, accountAndIDQueryCondition, accountID, serviceID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete service from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete service from store") - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "service %s not found", serviceID) - } - - return nil -} - -func (s *SqlStore) DeleteTarget(ctx context.Context, accountID string, serviceID string, targetID uint) error { - result := s.db.Delete(&rpservice.Target{}, "account_id = ? AND service_id = ? AND id = ?", accountID, serviceID, targetID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete target from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete target from store") - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "target not found for service %s", serviceID) - } - - return nil -} - -func (s *SqlStore) DeleteServiceTargets(ctx context.Context, accountID string, serviceID string) error { - result := s.db.Delete(&rpservice.Target{}, "account_id = ? AND service_id = ?", accountID, serviceID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete targets from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete targets from store") - } - - return nil -} - -// GetTargetsByServiceID retrieves all targets for a given service -func (s *SqlStore) GetTargetsByServiceID(ctx context.Context, lockStrength LockingStrength, accountID string, serviceID string) ([]*rpservice.Target, error) { - var targets []*rpservice.Target - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - result := tx.Where("account_id = ? AND service_id = ?", accountID, serviceID).Find(&targets) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get targets from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get targets from store") - } - - return targets, nil -} - -func (s *SqlStore) GetServiceByID(ctx context.Context, lockStrength LockingStrength, accountID, serviceID string) (*rpservice.Service, error) { - tx := s.db.Preload("Targets") - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var service *rpservice.Service - result := tx.Take(&service, accountAndIDQueryCondition, accountID, serviceID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "service %s not found", serviceID) - } - - log.WithContext(ctx).Errorf("failed to get service from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get service from store") - } - - if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt service data: %w", err) - } - - return service, nil -} - -func (s *SqlStore) GetServiceByDomain(ctx context.Context, domain string) (*rpservice.Service, error) { - var service *rpservice.Service - result := s.db.Preload("Targets").Where("domain = ?", domain).First(&service) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "service with domain %s not found", domain) - } - - log.WithContext(ctx).Errorf("failed to get service by domain from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get service by domain from store") - } - - if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt service data: %w", err) - } - - return service, nil -} - -func (s *SqlStore) GetServices(ctx context.Context, lockStrength LockingStrength) ([]*rpservice.Service, error) { - tx := s.db.Preload("Targets") - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var serviceList []*rpservice.Service - result := tx.Find(&serviceList) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get services from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get services from store") - } - - for _, service := range serviceList { - if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt service data: %w", err) - } - } - - return serviceList, nil -} - -func (s *SqlStore) GetAccountServices(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*rpservice.Service, error) { - tx := s.db.Preload("Targets") - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var serviceList []*rpservice.Service - result := tx.Find(&serviceList, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get services from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get services from store") - } - - for _, service := range serviceList { - if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { - return nil, fmt.Errorf("decrypt service data: %w", err) - } - } - - return serviceList, nil -} - -// RenewEphemeralService updates the last_renewed_at timestamp for an ephemeral service. -func (s *SqlStore) RenewEphemeralService(ctx context.Context, accountID, peerID, serviceID string) error { - result := s.db.Model(&rpservice.Service{}). - Where("id = ? AND account_id = ? AND source_peer = ? AND source = ?", serviceID, accountID, peerID, rpservice.SourceEphemeral). - Update("meta_last_renewed_at", time.Now()) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to renew ephemeral service: %v", result.Error) - return status.Errorf(status.Internal, "renew ephemeral service") - } - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "no active expose session for service %s", serviceID) - } - return nil -} - -// GetExpiredEphemeralServices returns ephemeral services whose last renewal exceeds the given TTL. -// Only the fields needed for reaping are selected. The limit parameter caps the batch size to -// avoid loading too many rows in a single tick. Rows with empty source_peer are excluded to -// skip malformed legacy data. -func (s *SqlStore) GetExpiredEphemeralServices(ctx context.Context, ttl time.Duration, limit int) ([]*rpservice.Service, error) { - cutoff := time.Now().Add(-ttl) - var services []*rpservice.Service - result := s.db. - Select("id", "account_id", "source_peer", "domain"). - Where("source = ? AND source_peer <> '' AND meta_last_renewed_at < ?", rpservice.SourceEphemeral, cutoff). - Limit(limit). - Find(&services) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get expired ephemeral services: %v", result.Error) - return nil, status.Errorf(status.Internal, "get expired ephemeral services") - } - return services, nil -} - -// CountEphemeralServicesByPeer returns the count of ephemeral services for a specific peer. -// Use LockingStrengthUpdate inside a transaction to serialize concurrent create operations. -// The locking is applied via a row-level SELECT ... FOR UPDATE (not on the aggregate) to -// stay compatible with Postgres, which disallows FOR UPDATE on COUNT(*). -func (s *SqlStore) CountEphemeralServicesByPeer(ctx context.Context, lockStrength LockingStrength, accountID, peerID string) (int64, error) { - if lockStrength == LockingStrengthNone { - var count int64 - result := s.db.Model(&rpservice.Service{}). - Where("account_id = ? AND source_peer = ? AND source = ?", accountID, peerID, rpservice.SourceEphemeral). - Count(&count) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to count ephemeral services: %v", result.Error) - return 0, status.Errorf(status.Internal, "count ephemeral services") - } - return count, nil - } - - var ids []string - result := s.db.Model(&rpservice.Service{}). - Clauses(clause.Locking{Strength: string(lockStrength)}). - Select("id"). - Where("account_id = ? AND source_peer = ? AND source = ?", accountID, peerID, rpservice.SourceEphemeral). - Pluck("id", &ids) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to count ephemeral services: %v", result.Error) - return 0, status.Errorf(status.Internal, "count ephemeral services") - } - return int64(len(ids)), nil -} - -// EphemeralServiceExists checks if an ephemeral service exists for the given peer and domain. -// Use LockingStrengthUpdate inside a transaction to serialize concurrent create operations. -func (s *SqlStore) EphemeralServiceExists(ctx context.Context, lockStrength LockingStrength, accountID, peerID, domain string) (bool, error) { - if lockStrength == LockingStrengthNone { - var count int64 - result := s.db.Model(&rpservice.Service{}). - Where("account_id = ? AND source_peer = ? AND domain = ? AND source = ?", accountID, peerID, domain, rpservice.SourceEphemeral). - Count(&count) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to check ephemeral service existence: %v", result.Error) - return false, status.Errorf(status.Internal, "check ephemeral service existence") - } - return count > 0, nil - } - - var id string - result := s.db.Model(&rpservice.Service{}). - Clauses(clause.Locking{Strength: string(lockStrength)}). - Select("id"). - Where("account_id = ? AND source_peer = ? AND domain = ? AND source = ?", accountID, peerID, domain, rpservice.SourceEphemeral). - Limit(1). - Pluck("id", &id) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to check ephemeral service existence: %v", result.Error) - return false, status.Errorf(status.Internal, "check ephemeral service existence") - } - return id != "", nil -} - -// GetServicesByClusterAndPort returns services matching the given proxy cluster, mode, and listen port. -func (s *SqlStore) GetServicesByClusterAndPort(ctx context.Context, lockStrength LockingStrength, proxyCluster string, mode string, listenPort uint16) ([]*rpservice.Service, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var services []*rpservice.Service - result := tx.Where("proxy_cluster = ? AND mode = ? AND listen_port = ?", proxyCluster, mode, listenPort).Find(&services) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "query services by cluster and port") - } - - return services, nil -} - -// GetServicesByCluster returns all services for the given proxy cluster. -func (s *SqlStore) GetServicesByCluster(ctx context.Context, lockStrength LockingStrength, proxyCluster string) ([]*rpservice.Service, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var services []*rpservice.Service - result := tx.Where("proxy_cluster = ?", proxyCluster).Find(&services) - if result.Error != nil { - return nil, status.Errorf(status.Internal, "query services by cluster") - } - return services, nil -} - -func (s *SqlStore) GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error) { - tx := s.db - - customDomain := &domain.Domain{} - result := tx.Take(&customDomain, accountAndIDQueryCondition, accountID, domainID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "custom domain %s not found", domainID) - } - - log.WithContext(ctx).Errorf("failed to get custom domain from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get custom domain from store") - } - - return customDomain, nil -} - -func (s *SqlStore) ListFreeDomains(ctx context.Context, accountID string) ([]string, error) { - return nil, nil -} - -func (s *SqlStore) ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) { - tx := s.db - - var domains []*domain.Domain - result := tx.Find(&domains, accountIDCondition, accountID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get reverse proxy custom domains from the store: %s", result.Error) - return nil, status.Errorf(status.Internal, "failed to get reverse proxy custom domains from store") - } - - return domains, nil -} - -// GetCustomDomainByName returns the custom domain row holding the given name, -// regardless of which account owns it. -func (s *SqlStore) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) { - customDomain := &domain.Domain{} - result := s.db.Take(customDomain, "domain = ?", domainName) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "custom domain %s not found", domainName) - } - - log.WithContext(ctx).Errorf("failed to get custom domain by name from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get custom domain from store") - } - - return customDomain, nil -} - -func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) { - newDomain := &domain.Domain{ - ID: xid.New().String(), // Generate our own ID because gorm doesn't always configure the database to handle this for us. - Domain: domainName, - AccountID: accountID, - TargetCluster: targetCluster, - Type: domain.TypeCustom, - Validated: validated, - } - if !validated { - expiresAt := time.Now().UTC().Add(domain.ValidationTTL) - newDomain.ValidationExpiresAt = &expiresAt - } - result := s.db.Create(newDomain) - if result.Error != nil { - // The unique index is the last guard when two requests clear the - // manager's availability check at the same time. The one that loses the - // insert is a conflict, not an internal failure. - var count int64 - if err := s.db.Model(&domain.Domain{}).Where("domain = ?", domainName).Count(&count).Error; err == nil && count > 0 { - // The insert error is logged even on this path: the name being taken - // is what the caller has to act on, but if the insert also failed for - // an unrelated reason the operator still needs to see it. - log.WithContext(ctx).Warnf("create reverse proxy custom domain %s rejected, name already registered: %v", domainName, result.Error) - return nil, status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName) - } - - log.WithContext(ctx).Errorf("failed to create reverse proxy custom domain to store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to create reverse proxy custom domain to store") - } - - return newDomain, nil -} - -// UpdateCustomDomain completes validation only while the original registration is pending. -func (s *SqlStore) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) { - if !d.Validated { - return nil, status.Errorf(status.InvalidArgument, "custom domain update must complete validation") - } - result := s.db.WithContext(ctx).Model(&domain.Domain{}). - Where(accountAndIDQueryCondition, accountID, d.ID). - Where("domain = ? AND target_cluster = ?", d.Domain, d.TargetCluster). - Where("validated = ? AND validation_expires_at > ?", false, time.Now().UTC()). - Update("validated", true) - if result.Error != nil { - return nil, fmt.Errorf("validate custom domain in store: %w", result.Error) - } - if result.RowsAffected == 0 { - return nil, status.Errorf(status.PreconditionFailed, "custom domain registration is no longer pending validation") - } - - return d, nil -} - -func (s *SqlStore) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error { - result := s.db.Delete(domain.Domain{}, accountAndIDQueryCondition, accountID, domainID) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete reverse proxy custom domain from store: %v", result.Error) - return status.Errorf(status.Internal, "failed to delete reverse proxy custom domain from store") - } - - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "reverse proxy custom domain %s not found", domainID) - } - - return nil -} - -// CreateAccessLog creates a new access log entry in the database -func (s *SqlStore) CreateAccessLog(ctx context.Context, logEntry *accesslogs.AccessLogEntry) error { - result := s.db.Create(logEntry) - if result.Error != nil { - log.WithContext(ctx).WithFields(log.Fields{ - "service_id": logEntry.ServiceID, - "method": logEntry.Method, - "host": logEntry.Host, - "path": logEntry.Path, - }).Errorf("failed to create access log entry in store: %v", result.Error) - return status.Errorf(status.Internal, "failed to create access log entry in store") - } - return nil -} - -// CreateAgentNetworkAccessLog persists a flattened agent-network access-log -// entry together with its authorising-group child rows in a single -// transaction. -func (s *SqlStore) CreateAgentNetworkAccessLog(ctx context.Context, entry *agentNetworkTypes.AgentNetworkAccessLog, groups []agentNetworkTypes.AgentNetworkAccessLogGroup) error { - err := s.db.Transaction(func(tx *gorm.DB) error { - // Idempotent on the log id / (log_id, group_id) so a proxy resend of the - // same entry can't fail the request. - if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(entry).Error; err != nil { - return err - } - if len(groups) > 0 { - if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&groups).Error; err != nil { - return err - } - } - return nil - }) - if err != nil { - log.WithContext(ctx).WithFields(log.Fields{ - "account_id": entry.AccountID, - "service_id": entry.ServiceID, - "model": entry.Model, - }).Errorf("failed to create agent-network access log entry in store: %v", err) - return status.Errorf(status.Internal, "failed to create agent-network access log entry in store") - } - return nil -} - -// CreateAgentNetworkUsage persists a stripped agent-network usage record -// together with its authorising-group child rows in a single transaction. -func (s *SqlStore) CreateAgentNetworkUsage(ctx context.Context, usage *agentNetworkTypes.AgentNetworkUsage, groups []agentNetworkTypes.AgentNetworkUsageGroup) error { - err := s.db.Transaction(func(tx *gorm.DB) error { - // Idempotent on the usage id / (usage_id, group_id) so a proxy resend of - // the same entry can't fail the request. - if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(usage).Error; err != nil { - return err - } - if len(groups) > 0 { - if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&groups).Error; err != nil { - return err - } - } - return nil - }) - if err != nil { - log.WithContext(ctx).WithFields(log.Fields{ - "account_id": usage.AccountID, - "model": usage.Model, - }).Errorf("failed to create agent-network usage record in store: %v", err) - return status.Errorf(status.Internal, "failed to create agent-network usage record in store") - } - return nil -} - -// DeleteOldAgentNetworkAccessLogs deletes an account's access-log rows (and -// their authorising-group child rows) older than the cutoff. Usage records are -// untouched — they are the long-term aggregate. Returns the number of log rows -// deleted. -func (s *SqlStore) DeleteOldAgentNetworkAccessLogs(ctx context.Context, accountID string, olderThan time.Time) (int64, error) { - var deleted int64 - err := s.db.Transaction(func(tx *gorm.DB) error { - // Remove group child rows for the soon-to-be-deleted logs first. - if err := tx.Exec( - "DELETE FROM agent_network_access_log_group WHERE account_id = ? AND log_id IN (SELECT id FROM agent_network_access_log WHERE account_id = ? AND timestamp < ?)", - accountID, accountID, olderThan, - ).Error; err != nil { - return err - } - res := tx.Where("account_id = ? AND timestamp < ?", accountID, olderThan). - Delete(&agentNetworkTypes.AgentNetworkAccessLog{}) - if res.Error != nil { - return res.Error - } - deleted = res.RowsAffected - return nil - }) - if err != nil { - log.WithContext(ctx).Errorf("failed to delete old agent-network access logs for account %s: %v", accountID, err) - return 0, status.Errorf(status.Internal, "failed to delete old agent-network access logs") - } - return deleted, nil -} - -// GetAgentNetworkUsageRows returns the stripped usage rows for an account that -// match the filter (date / user / group / provider / model). Aggregation into -// time buckets happens in the manager so granularities stay engine-portable. -func (s *SqlStore) GetAgentNetworkUsageRows(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkUsage, error) { - var rows []*agentNetworkTypes.AgentNetworkUsage - - query := s.applyAgentNetworkUsageFilters( - s.db.Where(accountIDCondition, accountID), - filter, - ).Order("timestamp ASC") - - if lockStrength != LockingStrengthNone { - query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - if err := query.Find(&rows).Error; err != nil { - log.WithContext(ctx).Errorf("failed to get agent-network usage rows from store: %v", err) - return nil, status.Errorf(status.Internal, "failed to get agent-network usage rows from store") - } - return rows, nil -} - -// applyAgentNetworkUsageFilters applies the shared access-log filter's -// date/user/group/provider/model conditions to a usage-table query. Pagination, -// sort and free-text search are ignored — the overview is an aggregate. -func (s *SqlStore) applyAgentNetworkUsageFilters(query *gorm.DB, filter agentNetworkTypes.AgentNetworkAccessLogFilter) *gorm.DB { - if filter.UserID != nil { - query = query.Where("user_id = ?", *filter.UserID) - } - if filter.SessionID != nil { - query = query.Where("session_id = ?", *filter.SessionID) - } - if len(filter.ProviderIDs) > 0 { - query = query.Where("resolved_provider_id IN ?", filter.ProviderIDs) - } - if len(filter.Models) > 0 { - query = query.Where("model IN ?", filter.Models) - } - if len(filter.GroupIDs) > 0 { - query = query.Where( - "id IN (SELECT usage_id FROM agent_network_request_usage_group WHERE group_id IN ?)", - filter.GroupIDs, - ) - } - if filter.StartDate != nil { - query = query.Where("timestamp >= ?", *filter.StartDate) - } - if filter.EndDate != nil { - query = query.Where("timestamp <= ?", *filter.EndDate) - } - return query -} - -// GetAgentNetworkAccessLogs retrieves flattened agent-network access logs for -// an account with server-side pagination, filtering and sorting. Authorising -// group ids are hydrated from the group child table for the returned page. -func (s *SqlStore) GetAgentNetworkAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLog, int64, error) { - var logs []*agentNetworkTypes.AgentNetworkAccessLog - var totalCount int64 - - countQuery := s.applyAgentNetworkAccessLogFilters( - s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}).Where(accountIDCondition, accountID), - filter, - ) - if err := countQuery.Count(&totalCount).Error; err != nil { - log.WithContext(ctx).Errorf("failed to count agent-network access logs: %v", err) - return nil, 0, status.Errorf(status.Internal, "failed to count agent-network access logs") - } - - query := s.applyAgentNetworkAccessLogFilters( - s.db.Where(accountIDCondition, accountID), - filter, - ). - Order(filter.GetSortColumn() + " " + filter.GetSortOrder()). - Limit(filter.GetLimit()). - Offset(filter.GetOffset()) - - if lockStrength != LockingStrengthNone { - query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - if err := query.Find(&logs).Error; err != nil { - log.WithContext(ctx).Errorf("failed to get agent-network access logs from store: %v", err) - return nil, 0, status.Errorf(status.Internal, "failed to get agent-network access logs from store") - } - - if err := s.hydrateAgentNetworkAccessLogGroups(ctx, accountID, logs); err != nil { - return nil, 0, err - } - - return logs, totalCount, nil -} - -// applyAgentNetworkAccessLogFilters applies the filter conditions to a query. -func (s *SqlStore) applyAgentNetworkAccessLogFilters(query *gorm.DB, filter agentNetworkTypes.AgentNetworkAccessLogFilter) *gorm.DB { - if filter.Search != nil { - p := "%" + *filter.Search + "%" - query = query.Where( - "id LIKE ? OR host LIKE ? OR path LIKE ? OR model LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)", - p, p, p, p, p, p, - ) - } - if filter.UserID != nil { - query = query.Where("user_id = ?", *filter.UserID) - } - if filter.SessionID != nil { - query = query.Where("session_id = ?", *filter.SessionID) - } - if filter.Decision != nil { - query = query.Where("decision = ?", *filter.Decision) - } - if filter.PathPrefix != nil { - query = query.Where("path LIKE ?", *filter.PathPrefix+"%") - } - if len(filter.ProviderIDs) > 0 { - query = query.Where("resolved_provider_id IN ?", filter.ProviderIDs) - } - if len(filter.Models) > 0 { - query = query.Where("model IN ?", filter.Models) - } - if len(filter.GroupIDs) > 0 { - query = query.Where( - "id IN (SELECT log_id FROM agent_network_access_log_group WHERE group_id IN ?)", - filter.GroupIDs, - ) - } - if filter.StartDate != nil { - query = query.Where("timestamp >= ?", *filter.StartDate) - } - if filter.EndDate != nil { - query = query.Where("timestamp <= ?", *filter.EndDate) - } - return query -} - -// hydrateAgentNetworkAccessLogGroups loads the authorising group ids for the -// given page of entries and assigns them onto each entry's GroupIDs field. -func (s *SqlStore) hydrateAgentNetworkAccessLogGroups(ctx context.Context, accountID string, logs []*agentNetworkTypes.AgentNetworkAccessLog) error { - if len(logs) == 0 { - return nil - } - - ids := make([]string, 0, len(logs)) - for _, l := range logs { - ids = append(ids, l.ID) - } - - var rows []agentNetworkTypes.AgentNetworkAccessLogGroup - if err := s.db. - Where(accountIDCondition, accountID). - Where("log_id IN ?", ids). - Find(&rows).Error; err != nil { - log.WithContext(ctx).Errorf("failed to hydrate agent-network access log groups: %v", err) - return status.Errorf(status.Internal, "failed to hydrate agent-network access log groups") - } - - byLog := make(map[string][]string, len(logs)) - for _, r := range rows { - byLog[r.LogID] = append(byLog[r.LogID], r.GroupID) - } - for _, l := range logs { - l.GroupIDs = byLog[l.ID] - } - return nil -} - -// agentNetworkSessionKeyExpr is the SQL group key for session-grouped access -// logs: the row's session id, or — when the client sent none — the row id, so -// session-less requests each form their own singleton group. COALESCE/NULLIF -// are standard SQL, so this stays portable across SQLite and Postgres. -const agentNetworkSessionKeyExpr = "COALESCE(NULLIF(session_id, ''), id)" - -// GetAgentNetworkAccessLogSessions retrieves agent-network access logs grouped -// by session, with server-side pagination, filtering and sorting at the session -// level. It paginates over the distinct session keys (ordered by the requested -// session-level aggregate), fetches every entry for the page's sessions, and -// folds them into per-session summaries. The returned count is the number of -// matching sessions. Filters apply to the entries, so a session's summary -// reflects only its filter-matching requests. -func (s *SqlStore) GetAgentNetworkAccessLogSessions(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLogSession, int64, error) { - // Count distinct sessions via a grouped subquery — portable and avoids - // relying on COUNT(DISTINCT ) quoting quirks. - sessionsSubquery := s.applyAgentNetworkAccessLogFilters( - s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}).Where(accountIDCondition, accountID), - filter, - ). - Select(agentNetworkSessionKeyExpr + " AS session_key"). - Group(agentNetworkSessionKeyExpr) - - var totalCount int64 - if err := s.db.Table("(?) AS sessions", sessionsSubquery).Count(&totalCount).Error; err != nil { - log.WithContext(ctx).Errorf("failed to count agent-network access-log sessions: %v", err) - return nil, 0, status.Errorf(status.Internal, "failed to count agent-network access-log sessions") - } - - // The page of session keys, ordered by the session-level aggregate. The - // session-key tiebreaker keeps pagination deterministic when the primary - // aggregate ties. - type sessionKeyRow struct { - SessionKey string - } - var keyRows []sessionKeyRow - keyQuery := s.applyAgentNetworkAccessLogFilters( - s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}).Where(accountIDCondition, accountID), - filter, - ). - Select(agentNetworkSessionKeyExpr + " AS session_key"). - Group(agentNetworkSessionKeyExpr). - Order(filter.GetSessionSortExpr() + " " + filter.GetSortOrder()). - Order("session_key ASC"). - Limit(filter.GetLimit()). - Offset(filter.GetOffset()) - if err := keyQuery.Scan(&keyRows).Error; err != nil { - log.WithContext(ctx).Errorf("failed to list agent-network access-log session keys: %v", err) - return nil, 0, status.Errorf(status.Internal, "failed to list agent-network access-log session keys") - } - if len(keyRows) == 0 { - return nil, totalCount, nil - } - - keys := make([]string, 0, len(keyRows)) - for _, r := range keyRows { - keys = append(keys, r.SessionKey) - } - - // All entries for the page's sessions, contiguous per session and oldest - // first within each — the fold relies on that ordering. - var entries []*agentNetworkTypes.AgentNetworkAccessLog - entriesQuery := s.applyAgentNetworkAccessLogFilters( - s.db.Where(accountIDCondition, accountID), - filter, - ). - Where(agentNetworkSessionKeyExpr+" IN ?", keys). - Order(agentNetworkSessionKeyExpr + ", timestamp ASC") - - if lockStrength != LockingStrengthNone { - entriesQuery = entriesQuery.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - if err := entriesQuery.Find(&entries).Error; err != nil { - log.WithContext(ctx).Errorf("failed to get agent-network access-log session entries: %v", err) - return nil, 0, status.Errorf(status.Internal, "failed to get agent-network access-log session entries") - } - - if err := s.hydrateAgentNetworkAccessLogGroups(ctx, accountID, entries); err != nil { - return nil, 0, err - } - - return agentNetworkTypes.FoldAccessLogSessions(keys, entries), totalCount, nil -} - -// GetAccountAccessLogs retrieves access logs for a given account with pagination and filtering -func (s *SqlStore) GetAccountAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) { - var logs []*accesslogs.AccessLogEntry - var totalCount int64 - - baseQuery := s.db. - Model(&accesslogs.AccessLogEntry{}). - Where(accountIDCondition, accountID) - - baseQuery = s.applyAccessLogFilters(baseQuery, filter) - - if err := baseQuery.Count(&totalCount).Error; err != nil { - log.WithContext(ctx).Errorf("failed to count access logs: %v", err) - return nil, 0, status.Errorf(status.Internal, "failed to count access logs") - } - - query := s.db. - Where(accountIDCondition, accountID) - - query = s.applyAccessLogFilters(query, filter) - - sortColumns := filter.GetSortColumn() - sortOrder := strings.ToUpper(filter.GetSortOrder()) - - var orderClauses []string - for _, col := range strings.Split(sortColumns, ",") { - col = strings.TrimSpace(col) - if col != "" { - orderClauses = append(orderClauses, col+" "+sortOrder) - } - } - orderClause := strings.Join(orderClauses, ", ") - - query = query. - Order(orderClause). - Limit(filter.GetLimit()). - Offset(filter.GetOffset()) - - if lockStrength != LockingStrengthNone { - query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - result := query.Find(&logs) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get access logs from store: %v", result.Error) - return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store") - } - - return logs, totalCount, nil -} - -// DeleteOldAccessLogs deletes all access logs older than the specified time -func (s *SqlStore) DeleteOldAccessLogs(ctx context.Context, olderThan time.Time) (int64, error) { - result := s.db. - Where("timestamp < ?", olderThan). - Delete(&accesslogs.AccessLogEntry{}) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error) - return 0, status.Errorf(status.Internal, "failed to delete old access logs") - } - - return result.RowsAffected, nil -} - -// applyAccessLogFilters applies filter conditions to the query -func (s *SqlStore) applyAccessLogFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB { - if filter.Search != nil { - searchPattern := "%" + *filter.Search + "%" - query = query.Where( - "id LIKE ? OR location_connection_ip LIKE ? OR host LIKE ? OR path LIKE ? OR CONCAT(host, path) LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)", - searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, - ) - } - - if filter.SourceIP != nil { - query = query.Where("location_connection_ip = ?", *filter.SourceIP) - } - - if filter.Host != nil { - query = query.Where("host = ?", *filter.Host) - } - - if filter.Path != nil { - // Support LIKE pattern for path filtering - query = query.Where("path LIKE ?", "%"+*filter.Path+"%") - } - - if filter.UserID != nil { - query = query.Where("user_id = ?", *filter.UserID) - } - - if filter.Method != nil { - query = query.Where("method = ?", *filter.Method) - } - - if filter.Status != nil { - switch *filter.Status { - case "success": - query = query.Where("status_code >= ? AND status_code < ?", 200, 400) - case "failed": - query = query.Where("status_code < ? OR status_code >= ?", 200, 400) - } - } - - if filter.StatusCode != nil { - query = query.Where("status_code = ?", *filter.StatusCode) - } - - if filter.StartDate != nil { - query = query.Where("timestamp >= ?", *filter.StartDate) - } - - if filter.EndDate != nil { - query = query.Where("timestamp <= ?", *filter.EndDate) - } - - return query -} - -func (s *SqlStore) GetServiceTargetByTargetID(ctx context.Context, lockStrength LockingStrength, accountID string, targetID string) (*rpservice.Target, error) { - tx := s.db - if lockStrength != LockingStrengthNone { - tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) - } - - var target *rpservice.Target - result := tx.Take(&target, "account_id = ? AND target_id = ?", accountID, targetID) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "service target with ID %s not found", targetID) - } - - log.WithContext(ctx).Errorf("failed to get service target from store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get service target from store") - } - - return target, nil -} - -// SaveProxy saves or updates a proxy in the database -func (s *SqlStore) SaveProxy(ctx context.Context, p *proxy.Proxy) error { - result := s.db.Save(p) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to save proxy: %v", result.Error) - return status.Errorf(status.Internal, "failed to save proxy") - } - return nil -} - -// DisconnectProxy marks a proxy as disconnected only if the session ID matches. -// This prevents a slow-to-close old session from overwriting a newer reconnection. -func (s *SqlStore) DisconnectProxy(ctx context.Context, proxyID, sessionID string) error { - now := time.Now() - result := s.db. - Model(&proxy.Proxy{}). - Where("id = ? AND session_id = ?", proxyID, sessionID). - Updates(map[string]any{ - "status": proxy.StatusDisconnected, - "disconnected_at": now, - "last_seen": now, - }) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to disconnect proxy %s session %s: %v", proxyID, sessionID, result.Error) - return status.Errorf(status.Internal, "failed to disconnect proxy") - } - if result.RowsAffected == 0 { - log.WithContext(ctx).Debugf("proxy %s session %s: no row updated (superseded by newer session)", proxyID, sessionID) - } - return nil -} - -// GetAllProxies returns all reverse proxy instance rows. -func (s *SqlStore) GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error) { - var proxies []*proxy.Proxy - result := s.db.Order("cluster_address, id").Find(&proxies) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get proxies: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get proxies") - } - return proxies, nil -} - -// DisconnectAllProxies force-marks every proxy that is not already disconnected -// as disconnected, regardless of session ID. Unlike DisconnectProxy it is not -// session-guarded: it is an administrative repair helper, not part of the -// connection lifecycle. last_seen is left untouched so the stale-proxy reaper -// keeps working off the real last heartbeat. Returns the number of proxies updated. -func (s *SqlStore) DisconnectAllProxies(ctx context.Context) (int64, error) { - result := s.db. - Model(&proxy.Proxy{}). - Where("status != ?", proxy.StatusDisconnected). - Updates(map[string]any{ - "status": proxy.StatusDisconnected, - "disconnected_at": time.Now(), - }) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to disconnect all proxies: %v", result.Error) - return 0, status.Errorf(status.Internal, "failed to disconnect all proxies") - } - return result.RowsAffected, nil -} - -// UpdateProxyHeartbeat updates the last_seen timestamp for the proxy's current session. -func (s *SqlStore) UpdateProxyHeartbeat(ctx context.Context, p *proxy.Proxy) error { - now := time.Now() - - result := s.db. - Model(&proxy.Proxy{}). - Where("id = ? AND session_id = ?", p.ID, p.SessionID). - Updates(map[string]any{ - "last_seen": now, - "status": proxy.StatusConnected, - "disconnected_at": nil, - }) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update proxy heartbeat: %v", result.Error) - return status.Errorf(status.Internal, "failed to update proxy heartbeat") - } - - if result.RowsAffected == 0 { - p.LastSeen = now - p.ConnectedAt = &now - p.Status = proxy.StatusConnected - if err := s.db.Create(p).Error; err != nil { - log.WithContext(ctx).Debugf("proxy %s session %s: heartbeat fallback insert skipped: %v", p.ID, p.SessionID, err) - } - } - - return nil -} - -// GetActiveProxyClusterAddresses returns the unique cluster addresses of active -// shared proxies (those without an account scope). BYOP cluster addresses are -// excluded; use GetActiveProxyClusterAddressesForAccount to retrieve them. -func (s *SqlStore) GetActiveProxyClusterAddresses(ctx context.Context) ([]string, error) { - var addresses []string - - result := s.db. - Model(&proxy.Proxy{}). - Where("account_id IS NULL AND status = ? AND last_seen > ?", proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). - Distinct("cluster_address"). - Pluck("cluster_address", &addresses) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get active proxy cluster addresses: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to get active proxy cluster addresses") - } - - return addresses, nil -} - -func (s *SqlStore) GetActiveProxyClusterAddressesForAccount(ctx context.Context, accountID string) ([]string, error) { - var addresses []string - - result := s.db. - Model(&proxy.Proxy{}). - Where("account_id = ? AND status = ? AND last_seen > ?", accountID, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). - Distinct("cluster_address"). - Pluck("cluster_address", &addresses) - - if result.Error != nil { - return nil, status.Errorf(status.Internal, "failed to get active proxy cluster addresses for account") - } - - return addresses, nil -} - -func (s *SqlStore) GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error) { - var p proxy.Proxy - result := s.db.Where("account_id = ?", accountID).Take(&p) - if result.Error != nil { - if errors.Is(result.Error, gorm.ErrRecordNotFound) { - return nil, status.Errorf(status.NotFound, "proxy not found for account") - } - return nil, status.Errorf(status.Internal, "get proxy by account ID: %v", result.Error) - } - return &p, nil -} - -func (s *SqlStore) CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error) { - var count int64 - result := s.db.Model(&proxy.Proxy{}).Where("account_id = ?", accountID).Count(&count) - if result.Error != nil { - return 0, status.Errorf(status.Internal, "count proxies by account ID: %v", result.Error) - } - return count, nil -} - -// HasActiveProxyAtClusterAddress reports whether any proxy — shared or -// account-scoped — is currently active at the given cluster address, using -// the same connected-within-threshold window as the other active-proxy -// queries. Backs the agent-network settings delete guard: settings cannot be -// deleted while a proxy declares the endpoint hostname as its address. -// -// The comparison folds case on both sides: the caller passes a normalized -// (lowercase) hostname, but proxies declare their cluster address verbatim -// and Connect stores it unchanged, so on case-sensitive collations a proxy -// declaring "GW.Example.com" would otherwise slip past the guard. Hostnames -// are case-insensitive per RFC 4343; the guard must be too. -func (s *SqlStore) HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error) { - var count int64 - result := s.db. - Model(&proxy.Proxy{}). - Where("LOWER(cluster_address) = LOWER(?) AND status = ? AND last_seen > ?", clusterAddress, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). - Count(&count) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to count active proxies at cluster address: %v", result.Error) - return false, status.Errorf(status.Internal, "failed to count active proxies at cluster address") - } - return count > 0, nil -} - -func (s *SqlStore) IsClusterAddressConflicting(ctx context.Context, clusterAddress, accountID string) (bool, error) { - var count int64 - result := s.db. - Model(&proxy.Proxy{}). - Where("cluster_address = ? AND (account_id IS NULL OR account_id != ?)", clusterAddress, accountID). - Count(&count) - if result.Error != nil { - return false, status.Errorf(status.Internal, "check cluster address conflict: %v", result.Error) - } - return count > 0, nil -} - -// HasForeignAccountProxyAtHost reports whether a proxy owned by a different -// account declares this host. Shared proxies (account_id IS NULL) are not -// foreign: a shared cluster is what most accounts pin their agent network -// gateway to. The match folds case because proxies declare their address as -// the operator spelled it while the caller's host is normalised; that costs a -// scan of the proxies table, taken once per account when its gateway is -// bootstrapped, not on the per-connect path IsClusterAddressConflicting serves. -func (s *SqlStore) HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) { - var count int64 - result := s.db. - Model(&proxy.Proxy{}). - Where("LOWER(cluster_address) = LOWER(?) AND account_id IS NOT NULL AND account_id != ?", host, accountID). - Count(&count) - if result.Error != nil { - return false, status.Errorf(status.Internal, "check proxy host ownership: %v", result.Error) - } - return count > 0, nil -} - -func (s *SqlStore) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error { - result := s.db. - Where("cluster_address = ? AND account_id = ?", clusterAddress, accountID). - Delete(&proxy.Proxy{}) - if result.Error != nil { - return status.Errorf(status.Internal, "delete account cluster: %v", result.Error) - } - if result.RowsAffected == 0 { - return status.Errorf(status.NotFound, "cluster not found") - } - return nil -} - -// GetProxyClusters returns every cluster the account can see (shared -// plus its own BYOP), regardless of whether any proxy in the cluster -// is currently heartbeating. Online and ConnectedProxies are derived -// from the 2-min active window so the dashboard can render offline -// clusters distinctly; the 1-hour heartbeat reaper still removes rows -// that go quiet for too long. -// -// AccountOwned is determined by whether any proxy row in the group -// carries a non-NULL account_id; the caller maps that to Cluster.Type. -// Capability flags are NOT filled here — the handler enriches them via -// the per-cluster capability lookups. -func (s *SqlStore) GetProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) { - activeCutoff := time.Now().Add(-proxyActiveThreshold) - - type clusterRow struct { - ID string - Address string - ConnectedProxies int - Online bool - AccountOwned bool - } - - var rows []clusterRow - result := s.db.Model(&proxy.Proxy{}). - Select( - "MIN(id) AS id, "+ - "cluster_address AS address, "+ - // COUNT(CASE WHEN ... THEN 1 END) counts only non-NULL — i.e. only - // rows that satisfy the predicate — so it works portably across - // sqlite/postgres/mysql without dialect-specific FILTER syntax. - "COUNT(CASE WHEN status = ? AND last_seen > ? THEN 1 END) AS connected_proxies, "+ - // MAX(CASE …) > 0 expresses BOOL_OR in a way Postgres tolerates - // (Postgres can't MAX a boolean column). - "MAX(CASE WHEN status = ? AND last_seen > ? THEN 1 ELSE 0 END) > 0 AS online, "+ - "MAX(CASE WHEN account_id IS NOT NULL THEN 1 ELSE 0 END) > 0 AS account_owned", - proxy.StatusConnected, activeCutoff, - proxy.StatusConnected, activeCutoff, - ). - Where("account_id IS NULL OR account_id = ?", accountID). - Group("cluster_address"). - Scan(&rows) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get proxy clusters: %v", result.Error) - return nil, status.Errorf(status.Internal, "get proxy clusters") - } - - clusters := make([]proxy.Cluster, 0, len(rows)) - for _, r := range rows { - c := proxy.Cluster{ - ID: r.ID, - Address: r.Address, - Online: r.Online, - ConnectedProxies: r.ConnectedProxies, - } - if r.AccountOwned { - c.Type = proxy.ClusterTypeAccount - } else { - c.Type = proxy.ClusterTypeShared - } - clusters = append(clusters, c) - } - - return clusters, nil -} - -// proxyActiveThreshold is the maximum age of a heartbeat for a proxy to be -// considered active. Must be at least 2x the heartbeat interval (1 min). -const proxyActiveThreshold = 2 * time.Minute - -var validCapabilityColumns = map[string]struct{}{ - "supports_custom_ports": {}, - "require_subdomain": {}, - "supports_crowdsec": {}, - "private": {}, -} - -// GetClusterSupportsCustomPorts returns whether any active proxy in the cluster -// supports custom ports. Returns nil when no proxy reported the capability. -func (s *SqlStore) GetClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool { - return s.getClusterCapability(ctx, clusterAddr, "supports_custom_ports") -} - -// GetClusterRequireSubdomain returns whether any active proxy in the cluster -// requires a subdomain. Returns nil when no proxy reported the capability. -func (s *SqlStore) GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { - return s.getClusterCapability(ctx, clusterAddr, "require_subdomain") -} - -// GetClusterSupportsPrivate reports whether any active proxy in the cluster -// has the private capability (nil = unreported). -func (s *SqlStore) GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool { - return s.getClusterCapability(ctx, clusterAddr, "private") -} - -// GetClusterSupportsCrowdSec returns whether all active proxies in the cluster -// have CrowdSec configured. Returns nil when no proxy reported the capability. -// Unlike other capabilities that use ANY-true (for rolling upgrades), CrowdSec -// requires unanimous support: a single unconfigured proxy would let requests -// bypass reputation checks. -func (s *SqlStore) GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool { - return s.getClusterUnanimousCapability(ctx, clusterAddr, "supports_crowdsec") -} - -// getClusterUnanimousCapability returns an aggregated boolean capability -// requiring all active proxies in the cluster to report true. -func (s *SqlStore) getClusterUnanimousCapability(ctx context.Context, clusterAddr, column string) *bool { - if _, ok := validCapabilityColumns[column]; !ok { - log.WithContext(ctx).Errorf("invalid capability column: %s", column) - return nil - } - - var result struct { - Total int64 - Reported int64 - AllTrue bool - } - - // All active proxies must have reported the capability (no NULLs) and all - // must report true. A single unreported or false proxy means the cluster - // does not unanimously support the capability. - err := s.db.WithContext(ctx). - Model(&proxy.Proxy{}). - Select("COUNT(*) AS total, "+ - "COUNT(CASE WHEN "+column+" IS NOT NULL THEN 1 END) AS reported, "+ - "COUNT(*) > 0 AND COUNT(*) = COUNT(CASE WHEN "+column+" = true THEN 1 END) AS all_true"). - Where("cluster_address = ? AND status = ? AND last_seen > ?", - clusterAddr, "connected", time.Now().Add(-proxyActiveThreshold)). - Scan(&result).Error - if err != nil { - log.WithContext(ctx).Errorf("query cluster capability %s for %s: %v", column, clusterAddr, err) - return nil - } - - if result.Total == 0 || result.Reported == 0 { - return nil - } - - // If any proxy has not reported (NULL), we can't confirm unanimous support. - if result.Reported < result.Total { - v := false - return &v - } - - return &result.AllTrue -} - -// getClusterCapability returns an aggregated boolean capability for the given -// cluster. It checks active (connected, recently seen) proxies and returns: -// - *true if any proxy in the cluster has the capability set to true, -// - *false if at least one proxy reported but none set it to true, -// - nil if no proxy reported the capability at all. -func (s *SqlStore) getClusterCapability(ctx context.Context, clusterAddr, column string) *bool { - if _, ok := validCapabilityColumns[column]; !ok { - log.WithContext(ctx).Errorf("invalid capability column: %s", column) - return nil - } - - var result struct { - HasCapability bool - AnyTrue bool - } - - err := s.db. - WithContext(ctx). - Model(&proxy.Proxy{}). - Select("COUNT(CASE WHEN "+column+" IS NOT NULL THEN 1 END) > 0 AS has_capability, "+ - "COALESCE(MAX(CASE WHEN "+column+" = true THEN 1 ELSE 0 END), 0) = 1 AS any_true"). - Where("cluster_address = ? AND status = ? AND last_seen > ?", - clusterAddr, "connected", time.Now().Add(-proxyActiveThreshold)). - Scan(&result).Error - if err != nil { - log.WithContext(ctx).Errorf("query cluster capability %s for %s: %v", column, clusterAddr, err) - return nil - } - - if !result.HasCapability { - return nil - } - - return &result.AnyTrue -} - -// CleanupStaleProxies deletes proxies that haven't sent heartbeat in the specified duration -func (s *SqlStore) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error { - cutoffTime := time.Now().Add(-inactivityDuration) - - result := s.db. - Where("last_seen < ?", cutoffTime). - Delete(&proxy.Proxy{}) - - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to cleanup stale proxies: %v", result.Error) - return status.Errorf(status.Internal, "failed to cleanup stale proxies") - } - - if result.RowsAffected > 0 { - log.WithContext(ctx).Infof("Cleaned up %d stale proxies", result.RowsAffected) - } - - return nil -} - -// GetRoutingPeerNetworks returns the distinct network names where the peer is assigned as a routing peer -// in an enabled network router, either directly or via peer groups. -func (s *SqlStore) GetRoutingPeerNetworks(_ context.Context, accountID, peerID string) ([]string, error) { - var routers []*routerTypes.NetworkRouter - if err := s.db.Select("peer, peer_groups, network_id").Where("account_id = ? AND enabled = true", accountID).Find(&routers).Error; err != nil { - return nil, status.Errorf(status.Internal, "failed to get enabled routers: %v", err) - } - - if len(routers) == 0 { - return nil, nil - } - - var groupPeers []types.GroupPeer - if err := s.db.Select("group_id").Where("account_id = ? AND peer_id = ?", accountID, peerID).Find(&groupPeers).Error; err != nil { - return nil, status.Errorf(status.Internal, "failed to get peer group memberships: %v", err) - } - - groupSet := make(map[string]struct{}, len(groupPeers)) - for _, gp := range groupPeers { - groupSet[gp.GroupID] = struct{}{} - } - - networkIDs := make(map[string]struct{}) - for _, r := range routers { - if r.Peer == peerID { - networkIDs[r.NetworkID] = struct{}{} - } else if r.Peer == "" { - for _, pg := range r.PeerGroups { - if _, ok := groupSet[pg]; ok { - networkIDs[r.NetworkID] = struct{}{} - break - } - } - } - } - - if len(networkIDs) == 0 { - return nil, nil - } - - ids := make([]string, 0, len(networkIDs)) - for id := range networkIDs { - ids = append(ids, id) - } - - var networks []*networkTypes.Network - if err := s.db.Select("name").Where("account_id = ? AND id IN ?", accountID, ids).Find(&networks).Error; err != nil { - return nil, status.Errorf(status.Internal, "failed to get networks: %v", err) - } - - names := make([]string, 0, len(networks)) - for _, n := range networks { - names = append(names, n.Name) - } - - return names, nil -} diff --git a/management/server/store/sql_store_account.go b/management/server/store/sql_store_account.go new file mode 100644 index 000000000..43fe09861 --- /dev/null +++ b/management/server/store/sql_store_account.go @@ -0,0 +1,1176 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "net" + "runtime/debug" + "strings" + "sync" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + nbdns "github.com/netbirdio/netbird/dns" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/status" +) + +// Deprecated: Full +// account operations are no longer supported +func (s *SqlStore) SaveAccount(ctx context.Context, account *types.Account) error { + start := time.Now() + defer func() { + elapsed := time.Since(start) + if elapsed > 1*time.Second { + log.WithContext(ctx).Tracef("SaveAccount for account %s exceeded 1s, took: %v", account.Id, elapsed) + } + }() + + // todo: remove this check after the issue is resolved + s.checkAccountDomainBeforeSave(ctx, account.Id, account.Domain) + + generateAccountSQLTypes(account) + + // Encrypt sensitive user data before saving + for i := range account.UsersG { + if err := account.UsersG[i].EncryptSensitiveData(s.fieldEncrypt); err != nil { + return fmt.Errorf("encrypt user: %w", err) + } + } + + for _, group := range account.GroupsG { + group.StoreGroupPeers() + } + + err := s.transaction(ctx, func(tx *gorm.DB) error { + result := tx.Select(clause.Associations).Delete(account.Policies, "account_id = ?", account.Id) + if result.Error != nil { + return result.Error + } + + result = tx.Select(clause.Associations).Delete(account.UsersG, "account_id = ?", account.Id) + if result.Error != nil { + return result.Error + } + + result = tx.Select(clause.Associations).Delete(account) + if result.Error != nil { + return result.Error + } + + result = tx. + Session(&gorm.Session{FullSaveAssociations: true}). + Clauses(clause.OnConflict{UpdateAll: true}). + Create(account) + if result.Error != nil { + return result.Error + } + return nil + }) + + took := time.Since(start) + if s.metrics != nil { + s.metrics.StoreMetrics().CountPersistenceDuration(took) + } + log.WithContext(ctx).Debugf("took %d ms to persist an account to the store", took.Milliseconds()) + + return err +} + +// generateAccountSQLTypes generates the GORM compatible types for the account +func generateAccountSQLTypes(account *types.Account) { + for _, key := range account.SetupKeys { + account.SetupKeysG = append(account.SetupKeysG, *key) + } + + if len(account.SetupKeys) != len(account.SetupKeysG) { + log.Warnf("SetupKeysG length mismatch for account %s", account.Id) + } + + for id, peer := range account.Peers { + peer.ID = id + account.PeersG = append(account.PeersG, *peer) + } + + for id, user := range account.Users { + user.Id = id + for id, pat := range user.PATs { + pat.ID = id + user.PATsG = append(user.PATsG, *pat) + } + account.UsersG = append(account.UsersG, *user) + } + + for id, group := range account.Groups { + group.ID = id + group.AccountID = account.Id + account.GroupsG = append(account.GroupsG, group) + } + + for id, route := range account.Routes { + route.ID = id + account.RoutesG = append(account.RoutesG, *route) + } + + for id, ns := range account.NameServerGroups { + ns.ID = id + account.NameServerGroupsG = append(account.NameServerGroupsG, *ns) + } +} + +// checkAccountDomainBeforeSave temporary method to troubleshoot an issue with domains getting blank +func (s *SqlStore) checkAccountDomainBeforeSave(ctx context.Context, accountID, newDomain string) { + var acc types.Account + var domain string + result := s.db.Model(&acc).Select("domain").Where(idQueryCondition, accountID).Take(&domain) + if result.Error != nil { + if !errors.Is(result.Error, gorm.ErrRecordNotFound) { + log.WithContext(ctx).Errorf("error when getting account %s from the store to check domain: %s", accountID, result.Error) + } + return + } + if domain != "" && newDomain == "" { + log.WithContext(ctx).Warnf("saving an account with empty domain when there was a domain set. Previous domain %s, Account ID: %s, Trace: %s", domain, accountID, debug.Stack()) + } +} + +func (s *SqlStore) DeleteAccount(ctx context.Context, account *types.Account) error { + start := time.Now() + + err := s.transaction(ctx, func(tx *gorm.DB) error { + result := tx.Select(clause.Associations).Delete(account.Policies, "account_id = ?", account.Id) + if result.Error != nil { + return result.Error + } + + result = tx.Select(clause.Associations).Delete(account.UsersG, "account_id = ?", account.Id) + if result.Error != nil { + return result.Error + } + + result = tx.Select(clause.Associations).Delete(account.Services, "account_id = ?", account.Id) + if result.Error != nil { + return result.Error + } + + if err := deleteAgentNetworkAccountConfig(tx, account.Id); err != nil { + return err + } + + result = tx.Select(clause.Associations).Delete(account) + if result.Error != nil { + return result.Error + } + + return nil + }) + + took := time.Since(start) + if s.metrics != nil { + s.metrics.StoreMetrics().CountPersistenceDuration(took) + } + log.WithContext(ctx).Tracef("took %d ms to delete an account to the store", took.Milliseconds()) + + return err +} + +// deleteAgentNetworkAccountConfig removes the account's agent network configuration. These +// tables are not account associations, so deleting the account does not reach them. The +// settings row holds the account's globally unique gateway domain and the provider rows +// hold its upstream API keys. Tables that grow with traffic are left out: consumption +// counters and access logs are swept in the background, and usage records are kept. +func deleteAgentNetworkAccountConfig(tx *gorm.DB, accountID string) error { + // Dependents first: policies point at providers and guardrails, and settings + // go last, as DeleteSettings refuses while providers exist. + models := []any{ + &agentNetworkTypes.Policy{}, + &agentNetworkTypes.Provider{}, + &agentNetworkTypes.Guardrail{}, + &agentNetworkTypes.AccountBudgetRule{}, + &agentNetworkTypes.Settings{}, + } + for _, model := range models { + if err := tx.Delete(model, "account_id = ?", accountID).Error; err != nil { + return fmt.Errorf("delete %T rows: %w", model, err) + } + } + return nil +} + +func (s *SqlStore) UpdateAccountDomainAttributes(ctx context.Context, accountID string, domain string, category string, isPrimaryDomain bool) error { + accountCopy := types.Account{ + Domain: domain, + DomainCategory: category, + IsDomainPrimaryAccount: isPrimaryDomain, + } + + fieldsToUpdate := []string{"domain", "domain_category", "is_domain_primary_account"} + result := s.db.Model(&types.Account{}). + Select(fieldsToUpdate). + Where(idQueryCondition, accountID). + Updates(&accountCopy) + if result.Error != nil { + return status.Errorf(status.Internal, "failed to update account domain attributes to store: %v", result.Error) + } + + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "account %s", accountID) + } + + return nil +} + +func (s *SqlStore) GetAccountByPrivateDomain(ctx context.Context, domain string) (*types.Account, error) { + accountID, err := s.GetAccountIDByPrivateDomain(ctx, LockingStrengthNone, domain) + if err != nil { + return nil, err + } + + // TODO: rework to not call GetAccount + return s.GetAccount(ctx, accountID) +} + +func (s *SqlStore) GetAccountIDByPrivateDomain(ctx context.Context, lockStrength LockingStrength, domain string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountID string + result := tx.Model(&types.Account{}).Select("id"). + Where("domain = ? and is_domain_primary_account = ? and domain_category = ?", + strings.ToLower(domain), true, types.PrivateCategory, + ).Take(&accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.Errorf(status.NotFound, "account not found: provided domain is not registered or is not private") + } + log.WithContext(ctx).Errorf("error when getting account from the store: %s", result.Error) + return "", status.NewGetAccountFromStoreError(result.Error) + } + + return accountID, nil +} + +func (s *SqlStore) GetAccountsCounter(ctx context.Context) (int64, error) { + var count int64 + result := s.db.Model(&types.Account{}).Count(&count) + if result.Error != nil { + return 0, fmt.Errorf("failed to get all accounts counter: %w", result.Error) + } + + return count, nil +} + +func (s *SqlStore) GetAllAccounts(ctx context.Context) (all []*types.Account) { + var accounts []types.Account + result := s.db.Find(&accounts) + if result.Error != nil { + return all + } + + for _, account := range accounts { + if acc, err := s.GetAccount(ctx, account.Id); err == nil { + all = append(all, acc) + } + } + + return all +} + +func (s *SqlStore) GetAccountMeta(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.AccountMeta, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountMeta types.AccountMeta + result := tx.Model(&types.Account{}). + Take(&accountMeta, idQueryCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("error when getting account meta %s from the store: %s", accountID, result.Error) + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewAccountNotFoundError(accountID) + } + return nil, status.NewGetAccountFromStoreError(result.Error) + } + + return &accountMeta, nil +} + +func (s *SqlStore) GetAccount(ctx context.Context, accountID string) (*types.Account, error) { + if s.pgxPool() != nil { + return s.getAccountPgx(ctx, accountID) + } + return s.getAccountGorm(ctx, accountID) +} + +func (s *SqlStore) getAccountGorm(ctx context.Context, accountID string) (*types.Account, error) { + start := time.Now() + defer func() { + elapsed := time.Since(start) + if elapsed > 1*time.Second { + log.WithContext(ctx).Tracef("GetAccount for account %s exceeded 1s, took: %v", accountID, elapsed) + } + }() + + var account types.Account + result := s.db.Model(&account). + Preload("UsersG.PATsG"). // have to be specified as this is nested reference + Preload("Policies.Rules"). + Preload("SetupKeysG"). + Preload("PeersG"). + Preload("UsersG"). + Preload("GroupsG.GroupPeers"). + Preload("RoutesG"). + Preload("NameServerGroupsG"). + Preload("PostureChecks"). + Preload("Networks"). + Preload("NetworkRouters"). + Preload("NetworkResources"). + Preload("Onboarding"). + Preload("Services.Targets"). + Preload("Domains"). + Take(&account, idQueryCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("error when getting account %s from the store: %s", accountID, result.Error) + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewAccountNotFoundError(accountID) + } + return nil, status.NewGetAccountFromStoreError(result.Error) + } + + account.SetupKeys = make(map[string]*types.SetupKey, len(account.SetupKeysG)) + for _, key := range account.SetupKeysG { + if key.UpdatedAt.IsZero() { + key.UpdatedAt = key.CreatedAt + } + if key.AutoGroups == nil { + key.AutoGroups = []string{} + } + account.SetupKeys[key.Key] = &key + } + account.SetupKeysG = nil + + account.Peers = make(map[string]*nbpeer.Peer, len(account.PeersG)) + for _, peer := range account.PeersG { + account.Peers[peer.ID] = &peer + } + account.PeersG = nil + account.Users = make(map[string]*types.User, len(account.UsersG)) + for _, user := range account.UsersG { + user.PATs = make(map[string]*types.PersonalAccessToken, len(user.PATs)) + for _, pat := range user.PATsG { + pat.UserID = "" + user.PATs[pat.ID] = &pat + } + if user.AutoGroups == nil { + user.AutoGroups = []string{} + } + if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt user: %w", err) + } + account.Users[user.Id] = &user + user.PATsG = nil + } + account.UsersG = nil + account.Groups = make(map[string]*types.Group, len(account.GroupsG)) + for _, group := range account.GroupsG { + group.Peers = make([]string, len(group.GroupPeers)) + for i, gp := range group.GroupPeers { + group.Peers[i] = gp.PeerID + } + if group.Resources == nil { + group.Resources = []types.Resource{} + } + account.Groups[group.ID] = group + } + account.GroupsG = nil + + account.Routes = make(map[route.ID]*route.Route, len(account.RoutesG)) + for _, route := range account.RoutesG { + account.Routes[route.ID] = &route + } + account.RoutesG = nil + account.NameServerGroups = make(map[string]*nbdns.NameServerGroup, len(account.NameServerGroupsG)) + for _, ns := range account.NameServerGroupsG { + ns.AccountID = "" + if ns.NameServers == nil { + ns.NameServers = []nbdns.NameServer{} + } + if ns.Groups == nil { + ns.Groups = []string{} + } + if ns.Domains == nil { + ns.Domains = []string{} + } + account.NameServerGroups[ns.ID] = &ns + } + account.NameServerGroupsG = nil + return &account, nil +} + +func (s *SqlStore) getAccountPgx(ctx context.Context, accountID string) (*types.Account, error) { + account, err := s.getAccount(ctx, accountID) + if err != nil { + return nil, err + } + + var wg sync.WaitGroup + errChan := make(chan error, 16) + + wg.Add(1) + go func() { + defer wg.Done() + keys, err := s.getSetupKeys(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.SetupKeysG = keys + }() + + wg.Add(1) + go func() { + defer wg.Done() + peers, err := s.getPeers(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.PeersG = peers + }() + + wg.Add(1) + go func() { + defer wg.Done() + users, err := s.getUsers(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.UsersG = users + }() + + wg.Add(1) + go func() { + defer wg.Done() + groups, err := s.getGroups(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.GroupsG = groups + }() + + wg.Add(1) + go func() { + defer wg.Done() + policies, err := s.getPolicies(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.Policies = policies + }() + + wg.Add(1) + go func() { + defer wg.Done() + routes, err := s.getRoutes(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.RoutesG = routes + }() + + wg.Add(1) + go func() { + defer wg.Done() + nsgs, err := s.getNameServerGroups(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.NameServerGroupsG = nsgs + }() + + wg.Add(1) + go func() { + defer wg.Done() + checks, err := s.getPostureChecks(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.PostureChecks = checks + }() + + wg.Add(1) + go func() { + defer wg.Done() + services, err := s.getServices(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.Services = services + }() + + wg.Add(1) + go func() { + defer wg.Done() + domains, err := s.ListCustomDomains(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.Domains = domains + }() + + wg.Add(1) + go func() { + defer wg.Done() + networks, err := s.getNetworks(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.Networks = networks + }() + + wg.Add(1) + go func() { + defer wg.Done() + routers, err := s.getNetworkRouters(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.NetworkRouters = routers + }() + + wg.Add(1) + go func() { + defer wg.Done() + resources, err := s.getNetworkResources(ctx, accountID) + if err != nil { + errChan <- err + return + } + account.NetworkResources = resources + }() + + wg.Add(1) + go func() { + defer wg.Done() + err := s.getAccountOnboarding(ctx, accountID, account) + if err != nil { + errChan <- err + return + } + }() + + wg.Wait() + close(errChan) + for e := range errChan { + if e != nil { + return nil, e + } + } + + var userIDs []string + for _, u := range account.UsersG { + userIDs = append(userIDs, u.Id) + } + var policyIDs []string + for _, p := range account.Policies { + policyIDs = append(policyIDs, p.ID) + } + var groupIDs []string + for _, g := range account.GroupsG { + groupIDs = append(groupIDs, g.ID) + } + + wg.Add(3) + errChan = make(chan error, 3) + + var pats []types.PersonalAccessToken + go func() { + defer wg.Done() + var err error + pats, err = s.getPersonalAccessTokens(ctx, userIDs) + if err != nil { + errChan <- err + } + }() + + var rules []*types.PolicyRule + go func() { + defer wg.Done() + var err error + rules, err = s.getPolicyRules(ctx, policyIDs) + if err != nil { + errChan <- err + } + }() + + var groupPeers []types.GroupPeer + go func() { + defer wg.Done() + var err error + groupPeers, err = s.getGroupPeers(ctx, groupIDs) + if err != nil { + errChan <- err + } + }() + + wg.Wait() + close(errChan) + for e := range errChan { + if e != nil { + return nil, e + } + } + + patsByUserID := make(map[string][]*types.PersonalAccessToken) + for i := range pats { + pat := &pats[i] + patsByUserID[pat.UserID] = append(patsByUserID[pat.UserID], pat) + pat.UserID = "" + } + + rulesByPolicyID := make(map[string][]*types.PolicyRule) + for _, rule := range rules { + rulesByPolicyID[rule.PolicyID] = append(rulesByPolicyID[rule.PolicyID], rule) + } + + peersByGroupID := make(map[string][]string) + for _, gp := range groupPeers { + peersByGroupID[gp.GroupID] = append(peersByGroupID[gp.GroupID], gp.PeerID) + } + + account.SetupKeys = make(map[string]*types.SetupKey, len(account.SetupKeysG)) + for i := range account.SetupKeysG { + key := &account.SetupKeysG[i] + account.SetupKeys[key.Key] = key + } + + account.Peers = make(map[string]*nbpeer.Peer, len(account.PeersG)) + for i := range account.PeersG { + peer := &account.PeersG[i] + account.Peers[peer.ID] = peer + } + + account.Users = make(map[string]*types.User, len(account.UsersG)) + for i := range account.UsersG { + user := &account.UsersG[i] + if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt user: %w", err) + } + user.PATs = make(map[string]*types.PersonalAccessToken) + if userPats, ok := patsByUserID[user.Id]; ok { + for j := range userPats { + pat := userPats[j] + user.PATs[pat.ID] = pat + } + } + account.Users[user.Id] = user + } + + for i := range account.Policies { + policy := account.Policies[i] + if policyRules, ok := rulesByPolicyID[policy.ID]; ok { + policy.Rules = policyRules + } + } + + account.Groups = make(map[string]*types.Group, len(account.GroupsG)) + for i := range account.GroupsG { + group := account.GroupsG[i] + if peerIDs, ok := peersByGroupID[group.ID]; ok { + group.Peers = peerIDs + } + account.Groups[group.ID] = group + } + + account.Routes = make(map[route.ID]*route.Route, len(account.RoutesG)) + for i := range account.RoutesG { + route := &account.RoutesG[i] + account.Routes[route.ID] = route + } + + account.NameServerGroups = make(map[string]*nbdns.NameServerGroup, len(account.NameServerGroupsG)) + for i := range account.NameServerGroupsG { + nsg := &account.NameServerGroupsG[i] + nsg.AccountID = "" + account.NameServerGroups[nsg.ID] = nsg + } + + account.SetupKeysG = nil + account.PeersG = nil + account.UsersG = nil + account.GroupsG = nil + account.RoutesG = nil + account.NameServerGroupsG = nil + + return account, nil +} + +func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Account, error) { + var account types.Account + account.Network = &types.Network{} + const accountQuery = ` + SELECT + id, created_by, created_at, domain, domain_category, is_domain_primary_account, + -- Embedded Network + network_identifier, network_net, network_net_v6, network_dns, network_serial, + -- Embedded DNSSettings + dns_settings_disabled_management_groups, + -- Embedded Settings + settings_peer_login_expiration_enabled, settings_peer_login_expiration, + settings_peer_inactivity_expiration_enabled, settings_peer_inactivity_expiration, + settings_regular_users_view_blocked, settings_groups_propagation_enabled, + settings_jwt_groups_enabled, settings_jwt_groups_claim_name, settings_jwt_allow_groups, + settings_routing_peer_dns_resolution_enabled, settings_dns_domain, settings_network_range, + settings_network_range_v6, settings_ipv6_enabled_groups, settings_lazy_connection_enabled, + settings_local_mfa_enabled, settings_metrics_push_enabled, settings_agent_network_only, + settings_dashboard_features, settings_auto_update_version, settings_auto_update_always, + settings_peer_expose_enabled, settings_peer_expose_groups, + -- Embedded ExtraSettings + settings_extra_peer_approval_enabled, settings_extra_user_approval_required, + settings_extra_integrated_validator, settings_extra_integrated_validator_groups + FROM accounts WHERE id = $1` + + var ( + sPeerLoginExpirationEnabled sql.NullBool + sPeerLoginExpiration sql.NullInt64 + sPeerInactivityExpirationEnabled sql.NullBool + sPeerInactivityExpiration sql.NullInt64 + sRegularUsersViewBlocked sql.NullBool + sGroupsPropagationEnabled sql.NullBool + sJWTGroupsEnabled sql.NullBool + sJWTGroupsClaimName sql.NullString + sJWTAllowGroups sql.NullString + sRoutingPeerDNSResolutionEnabled sql.NullBool + sDNSDomain sql.NullString + sNetworkRange sql.NullString + sNetworkRangeV6 sql.NullString + sIPv6EnabledGroups sql.NullString + sLazyConnectionEnabled sql.NullBool + sLocalMFAEnabled sql.NullBool + sMetricsPushEnabled sql.NullBool + sAgentNetworkOnly sql.NullBool + sDashboardFeatures sql.NullString + autoUpdateVersion sql.NullString + autoUpdateAlways sql.NullBool + peerExposeEnabled sql.NullBool + peerExposeGroups sql.NullString + sExtraPeerApprovalEnabled sql.NullBool + sExtraUserApprovalRequired sql.NullBool + sExtraIntegratedValidator sql.NullString + sExtraIntegratedValidatorGroups sql.NullString + networkNet sql.NullString + networkNetV6 sql.NullString + dnsSettingsDisabledGroups sql.NullString + networkIdentifier sql.NullString + networkDns sql.NullString + networkSerial sql.NullInt64 + createdAt sql.NullTime + ) + err := s.pgxPool().QueryRow(ctx, accountQuery, accountID).Scan( + &account.Id, &account.CreatedBy, &createdAt, &account.Domain, &account.DomainCategory, &account.IsDomainPrimaryAccount, + &networkIdentifier, &networkNet, &networkNetV6, &networkDns, &networkSerial, + &dnsSettingsDisabledGroups, + &sPeerLoginExpirationEnabled, &sPeerLoginExpiration, + &sPeerInactivityExpirationEnabled, &sPeerInactivityExpiration, + &sRegularUsersViewBlocked, &sGroupsPropagationEnabled, + &sJWTGroupsEnabled, &sJWTGroupsClaimName, &sJWTAllowGroups, + &sRoutingPeerDNSResolutionEnabled, &sDNSDomain, &sNetworkRange, + &sNetworkRangeV6, &sIPv6EnabledGroups, &sLazyConnectionEnabled, + &sLocalMFAEnabled, &sMetricsPushEnabled, &sAgentNetworkOnly, + &sDashboardFeatures, &autoUpdateVersion, &autoUpdateAlways, + &peerExposeEnabled, &peerExposeGroups, + &sExtraPeerApprovalEnabled, &sExtraUserApprovalRequired, + &sExtraIntegratedValidator, &sExtraIntegratedValidatorGroups, + ) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, status.NewAccountNotFoundError(accountID) + } + return nil, status.NewGetAccountFromStoreError(err) + } + + account.Settings = &types.Settings{Extra: &types.ExtraSettings{}} + if networkNet.Valid { + _ = json.Unmarshal([]byte(networkNet.String), &account.Network.Net) + } + if createdAt.Valid { + account.CreatedAt = createdAt.Time + } + if dnsSettingsDisabledGroups.Valid { + _ = json.Unmarshal([]byte(dnsSettingsDisabledGroups.String), &account.DNSSettings.DisabledManagementGroups) + } + if networkIdentifier.Valid { + account.Network.Identifier = networkIdentifier.String + } + if networkDns.Valid { + account.Network.Dns = networkDns.String + } + if networkSerial.Valid { + account.Network.Serial = uint64(networkSerial.Int64) + } + if sPeerLoginExpirationEnabled.Valid { + account.Settings.PeerLoginExpirationEnabled = sPeerLoginExpirationEnabled.Bool + } + if sPeerLoginExpiration.Valid { + account.Settings.PeerLoginExpiration = time.Duration(sPeerLoginExpiration.Int64) + } + if sPeerInactivityExpirationEnabled.Valid { + account.Settings.PeerInactivityExpirationEnabled = sPeerInactivityExpirationEnabled.Bool + } + if sPeerInactivityExpiration.Valid { + account.Settings.PeerInactivityExpiration = time.Duration(sPeerInactivityExpiration.Int64) + } + if sRegularUsersViewBlocked.Valid { + account.Settings.RegularUsersViewBlocked = sRegularUsersViewBlocked.Bool + } + if sGroupsPropagationEnabled.Valid { + account.Settings.GroupsPropagationEnabled = sGroupsPropagationEnabled.Bool + } + if sJWTGroupsEnabled.Valid { + account.Settings.JWTGroupsEnabled = sJWTGroupsEnabled.Bool + } + if sJWTGroupsClaimName.Valid { + account.Settings.JWTGroupsClaimName = sJWTGroupsClaimName.String + } + if sRoutingPeerDNSResolutionEnabled.Valid { + account.Settings.RoutingPeerDNSResolutionEnabled = sRoutingPeerDNSResolutionEnabled.Bool + } + if sDNSDomain.Valid { + account.Settings.DNSDomain = sDNSDomain.String + } + if sLazyConnectionEnabled.Valid { + account.Settings.LazyConnectionEnabled = sLazyConnectionEnabled.Bool + } + if sLocalMFAEnabled.Valid { + account.Settings.LocalMfaEnabled = sLocalMFAEnabled.Bool + } + if sMetricsPushEnabled.Valid { + account.Settings.MetricsPushEnabled = sMetricsPushEnabled.Bool + } + if sAgentNetworkOnly.Valid { + account.Settings.AgentNetworkOnly = sAgentNetworkOnly.Bool + } + if sDashboardFeatures.Valid && sDashboardFeatures.String != "" { + if err := json.Unmarshal([]byte(sDashboardFeatures.String), &account.Settings.DashboardFeatures); err != nil { + log.WithContext(ctx).Warnf("failed to unmarshal dashboard features for account %s: %v", accountID, err) + } + } + if sJWTAllowGroups.Valid { + _ = json.Unmarshal([]byte(sJWTAllowGroups.String), &account.Settings.JWTAllowGroups) + } + if sNetworkRange.Valid { + _ = json.Unmarshal([]byte(sNetworkRange.String), &account.Settings.NetworkRange) + } + if networkNetV6.Valid { + _ = json.Unmarshal([]byte(networkNetV6.String), &account.Network.NetV6) + } + if sNetworkRangeV6.Valid { + _ = json.Unmarshal([]byte(sNetworkRangeV6.String), &account.Settings.NetworkRangeV6) + } + if sIPv6EnabledGroups.Valid { + _ = json.Unmarshal([]byte(sIPv6EnabledGroups.String), &account.Settings.IPv6EnabledGroups) + } + if autoUpdateAlways.Valid { + account.Settings.AutoUpdateAlways = autoUpdateAlways.Bool + } + if autoUpdateVersion.Valid { + account.Settings.AutoUpdateVersion = autoUpdateVersion.String + } + if peerExposeEnabled.Valid { + account.Settings.PeerExposeEnabled = peerExposeEnabled.Bool + } + if peerExposeGroups.Valid { + _ = json.Unmarshal([]byte(peerExposeGroups.String), &account.Settings.PeerExposeGroups) + } + + if sExtraPeerApprovalEnabled.Valid { + account.Settings.Extra.PeerApprovalEnabled = sExtraPeerApprovalEnabled.Bool + } + if sExtraUserApprovalRequired.Valid { + account.Settings.Extra.UserApprovalRequired = sExtraUserApprovalRequired.Bool + } + if sExtraIntegratedValidator.Valid { + account.Settings.Extra.IntegratedValidator = sExtraIntegratedValidator.String + } + if sExtraIntegratedValidatorGroups.Valid { + _ = json.Unmarshal([]byte(sExtraIntegratedValidatorGroups.String), &account.Settings.Extra.IntegratedValidatorGroups) + } + return &account, nil +} + +func (s *SqlStore) GetAnyAccountID(ctx context.Context) (string, error) { + var account types.Account + result := s.db.Select("id").Order("created_at desc").Limit(1).Find(&account) + if result.Error != nil { + return "", status.NewGetAccountFromStoreError(result.Error) + } + if result.RowsAffected == 0 { + return "", status.Errorf(status.NotFound, "account not found: index lookup failed") + } + + return account.Id, nil +} + +func (s *SqlStore) GetAccountNetwork(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.Network, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountNetwork types.AccountNetwork + if err := tx.Model(&types.Account{}).Where(idQueryCondition, accountID).Take(&accountNetwork).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewAccountNotFoundError(accountID) + } + return nil, status.Errorf(status.Internal, "issue getting network from store: %s", err) + } + return accountNetwork.Network, nil +} + +func (s *SqlStore) GetAccountSettings(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.Settings, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountSettings types.AccountSettings + if err := tx.Model(&types.Account{}).Where(idQueryCondition, accountID).Take(&accountSettings).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "settings not found") + } + return nil, status.Errorf(status.Internal, "issue getting settings from store: %s", err) + } + return accountSettings.Settings, nil +} + +func (s *SqlStore) GetAccountCreatedBy(ctx context.Context, lockStrength LockingStrength, accountID string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var createdBy string + result := tx.Model(&types.Account{}). + Select("created_by").Take(&createdBy, idQueryCondition, accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.NewAccountNotFoundError(accountID) + } + return "", status.NewGetAccountFromStoreError(result.Error) + } + + return createdBy, nil +} + +func (s *SqlStore) IncrementNetworkSerial(ctx context.Context, accountId string) error { + result := s.db.Model(&types.Account{}).Where(idQueryCondition, accountId).Update("network_serial", gorm.Expr("network_serial + 1")) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to increment network serial count in store: %v", result.Error) + return status.Errorf(status.Internal, "failed to increment network serial count in store") + } + return nil +} + +func (s *SqlStore) GetAccountDNSSettings(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.DNSSettings, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountDNSSettings types.AccountDNSSettings + result := tx.Model(&types.Account{}). + Take(&accountDNSSettings, idQueryCondition, accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewAccountNotFoundError(accountID) + } + log.WithContext(ctx).Errorf("failed to get dns settings from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get dns settings from store") + } + return &accountDNSSettings.DNSSettings, nil +} + +// AccountExists checks whether an account exists by the given ID. +func (s *SqlStore) AccountExists(ctx context.Context, lockStrength LockingStrength, id string) (bool, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountID string + result := tx.Model(&types.Account{}). + Select("id").Take(&accountID, idQueryCondition, id) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return false, nil + } + return false, result.Error + } + + return accountID != "", nil +} + +// GetAccountDomainAndCategory retrieves the Domain and DomainCategory fields for an account based on the given accountID. +func (s *SqlStore) GetAccountDomainAndCategory(ctx context.Context, lockStrength LockingStrength, accountID string) (string, string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var account types.Account + result := tx.Model(&types.Account{}).Select("domain", "domain_category"). + Where(idQueryCondition, accountID).Take(&account) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", "", status.Errorf(status.NotFound, "account not found") + } + return "", "", status.Errorf(status.Internal, "failed to get domain category from store: %v", result.Error) + } + + return account.Domain, account.DomainCategory, nil +} + +// SaveDNSSettings saves the DNS settings to the store. +func (s *SqlStore) SaveDNSSettings(ctx context.Context, accountID string, settings *types.DNSSettings) error { + result := s.db.Model(&types.Account{}). + Where(idQueryCondition, accountID).Updates(&types.AccountDNSSettings{DNSSettings: *settings}) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save dns settings to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to save dns settings to store") + } + + if result.RowsAffected == 0 { + return status.NewAccountNotFoundError(accountID) + } + + return nil +} + +// SaveAccountSettings stores the account settings in DB. +func (s *SqlStore) SaveAccountSettings(ctx context.Context, accountID string, settings *types.Settings) error { + result := s.db.Model(&types.Account{}). + Select("*").Where(idQueryCondition, accountID).Updates(&types.AccountSettings{Settings: settings}) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save account settings to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to save account settings to store") + } + + // MySQL reports RowsAffected=0 for no-op updates where values don't change, + // unlike SQLite/Postgres which report matched rows. Skip the check since the + // caller (UpdateAccountSettings) already verified the account exists via + // GetAccountSettings with LockingStrengthUpdate. + + return nil +} + +func (s *SqlStore) CountAccountsByPrivateDomain(ctx context.Context, domain string) (int64, error) { + var count int64 + result := s.db.Model(&types.Account{}). + Where("domain = ? AND domain_category = ?", + strings.ToLower(domain), types.PrivateCategory, + ).Count(&count) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to count accounts by private domain %s: %s", domain, result.Error) + return 0, status.Errorf(status.Internal, "failed to count accounts by private domain") + } + + return count, nil +} + +func (s *SqlStore) IsPrimaryAccount(ctx context.Context, accountID string) (bool, string, error) { + var info types.PrimaryAccountInfo + result := s.db.Model(&types.Account{}). + Select("is_domain_primary_account, domain"). + Where(idQueryCondition, accountID). + Take(&info) + + if result.Error != nil { + return false, "", status.Errorf(status.Internal, "failed to get account info: %v", result.Error) + } + + return info.IsDomainPrimaryAccount, info.Domain, nil +} + +func (s *SqlStore) MarkAccountPrimary(ctx context.Context, accountID string) error { + result := s.db.Model(&types.Account{}). + Where(idQueryCondition, accountID). + Update("is_domain_primary_account", true) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to mark account as primary: %s", result.Error) + return status.Errorf(status.Internal, "failed to mark account as primary") + } + + if result.RowsAffected == 0 { + return status.NewAccountNotFoundError(accountID) + } + + return nil +} + +type accountNetworkPatch struct { + Network *types.Network `gorm:"embedded;embeddedPrefix:network_"` +} + +func (s *SqlStore) UpdateAccountNetwork(ctx context.Context, accountID string, ipNet net.IPNet) error { + patch := accountNetworkPatch{ + Network: &types.Network{Net: ipNet}, + } + + result := s.db. + Model(&types.Account{}). + Where(idQueryCondition, accountID). + Updates(&patch) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update account network: %v", result.Error) + return status.Errorf(status.Internal, "failed to update account network") + } + if result.RowsAffected == 0 { + return status.NewAccountNotFoundError(accountID) + } + return nil +} + +// UpdateAccountNetworkV6 updates the IPv6 network range for the account. +func (s *SqlStore) UpdateAccountNetworkV6(ctx context.Context, accountID string, ipNet net.IPNet) error { + patch := accountNetworkPatch{ + Network: &types.Network{NetV6: ipNet}, + } + + result := s.db. + Model(&types.Account{}). + Where(idQueryCondition, accountID). + Updates(&patch) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update account network v6: %v", result.Error) + return status.Errorf(status.Internal, "update account network v6") + } + if result.RowsAffected == 0 { + return status.NewAccountNotFoundError(accountID) + } + return nil +} diff --git a/management/server/store/sql_store_account_onboarding.go b/management/server/store/sql_store_account_onboarding.go new file mode 100644 index 000000000..73c8f14d0 --- /dev/null +++ b/management/server/store/sql_store_account_onboarding.go @@ -0,0 +1,70 @@ +package store + +import ( + "context" + "database/sql" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// GetAccountOnboarding retrieves the onboarding information for a specific account. +func (s *SqlStore) GetAccountOnboarding(ctx context.Context, accountID string) (*types.AccountOnboarding, error) { + var accountOnboarding types.AccountOnboarding + result := s.db.Model(&accountOnboarding).Take(&accountOnboarding, accountIDCondition, accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewAccountOnboardingNotFoundError(accountID) + } + log.WithContext(ctx).Errorf("error when getting account onboarding %s from the store: %s", accountID, result.Error) + return nil, status.NewGetAccountFromStoreError(result.Error) + } + + return &accountOnboarding, nil +} + +// SaveAccountOnboarding updates the onboarding information for a specific account. +func (s *SqlStore) SaveAccountOnboarding(ctx context.Context, onboarding *types.AccountOnboarding) error { + result := s.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(onboarding) + if result.Error != nil { + log.WithContext(ctx).Errorf("error when saving account onboarding %s in the store: %s", onboarding.AccountID, result.Error) + return status.Errorf(status.Internal, "error when saving account onboarding %s in the store: %s", onboarding.AccountID, result.Error) + } + + return nil +} + +func (s *SqlStore) getAccountOnboarding(ctx context.Context, accountID string, account *types.Account) error { + const query = `SELECT account_id, onboarding_flow_pending, signup_form_pending, created_at, updated_at FROM account_onboardings WHERE account_id = $1` + var onboardingFlowPending, signupFormPending sql.NullBool + var createdAt, updatedAt sql.NullTime + err := s.pgxPool().QueryRow(ctx, query, accountID).Scan( + &account.Onboarding.AccountID, + &onboardingFlowPending, + &signupFormPending, + &createdAt, + &updatedAt, + ) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return err + } + if createdAt.Valid { + account.Onboarding.CreatedAt = createdAt.Time + } + if updatedAt.Valid { + account.Onboarding.UpdatedAt = updatedAt.Time + } + if onboardingFlowPending.Valid { + account.Onboarding.OnboardingFlowPending = onboardingFlowPending.Bool + } + if signupFormPending.Valid { + account.Onboarding.SignupFormPending = signupFormPending.Bool + } + return nil +} diff --git a/management/server/store/sql_store_account_onboarding_test.go b/management/server/store/sql_store_account_onboarding_test.go new file mode 100644 index 000000000..9531cd892 --- /dev/null +++ b/management/server/store/sql_store_account_onboarding_test.go @@ -0,0 +1,68 @@ +package store + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/types" +) + +func TestSqlStore_GetAccountOnboarding(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "9439-34653001fc3b-bf1c8084-ba50-4ce7" + a, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + t.Logf("Onboarding: %+v", a.Onboarding) + err = store.SaveAccount(context.Background(), a) + require.NoError(t, err) + onboarding, err := store.GetAccountOnboarding(context.Background(), accountID) + require.NoError(t, err) + require.NotNil(t, onboarding) + require.Equal(t, accountID, onboarding.AccountID) + require.Equal(t, time.Date(2024, time.October, 2, 14, 1, 38, 210000000, time.UTC), onboarding.CreatedAt.UTC()) +} + +func TestSqlStore_SaveAccountOnboarding(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + t.Run("New onboarding should be saved correctly", func(t *testing.T) { + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + onboarding := &types.AccountOnboarding{ + AccountID: accountID, + SignupFormPending: true, + OnboardingFlowPending: true, + } + + err = store.SaveAccountOnboarding(context.Background(), onboarding) + require.NoError(t, err) + + savedOnboarding, err := store.GetAccountOnboarding(context.Background(), accountID) + require.NoError(t, err) + require.Equal(t, onboarding.SignupFormPending, savedOnboarding.SignupFormPending) + require.Equal(t, onboarding.OnboardingFlowPending, savedOnboarding.OnboardingFlowPending) + }) + + t.Run("Existing onboarding should be updated correctly", func(t *testing.T) { + accountID := "9439-34653001fc3b-bf1c8084-ba50-4ce7" + onboarding, err := store.GetAccountOnboarding(context.Background(), accountID) + require.NoError(t, err) + + onboarding.OnboardingFlowPending = !onboarding.OnboardingFlowPending + onboarding.SignupFormPending = !onboarding.SignupFormPending + + err = store.SaveAccountOnboarding(context.Background(), onboarding) + require.NoError(t, err) + + savedOnboarding, err := store.GetAccountOnboarding(context.Background(), accountID) + require.NoError(t, err) + require.Equal(t, onboarding.SignupFormPending, savedOnboarding.SignupFormPending) + require.Equal(t, onboarding.OnboardingFlowPending, savedOnboarding.OnboardingFlowPending) + }) +} diff --git a/management/server/store/sql_store_account_test.go b/management/server/store/sql_store_account_test.go new file mode 100644 index 000000000..ec5976982 --- /dev/null +++ b/management/server/store/sql_store_account_test.go @@ -0,0 +1,1011 @@ +package store + +import ( + "context" + "encoding/binary" + "fmt" + "net" + "net/netip" + "os" + "reflect" + "runtime" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + nbdns "github.com/netbirdio/netbird/dns" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + proxydomain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + networkTypes "github.com/netbirdio/netbird/management/server/networks/types" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" + nbroute "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/status" + "github.com/netbirdio/netbird/shared/testing_helpers" +) + +func Test_SaveAccount_Large(t *testing.T) { + if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { + t.Skip("skip CI tests on darwin and windows") + } + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + runLargeTest(t, store) + }) +} + +func runLargeTest(t *testing.T, store Store) { + t.Helper() + + account := newAccountWithId(context.Background(), "account_id", "testuser", "") + groupALL, err := account.GetGroupAll() + if err != nil { + t.Fatal(err) + } + setupKey, _ := types.GenerateDefaultSetupKey() + account.SetupKeys[setupKey.Key] = setupKey + const numPerAccount = 6000 + for n := 0; n < numPerAccount; n++ { + netIP := sequentialIPv4(n) + peerID := fmt.Sprintf("%s-peer-%d", account.Id, n) + addr, _ := netip.AddrFromSlice(netIP) + + peer := &nbpeer.Peer{ + ID: peerID, + Key: peerID, + IP: addr.Unmap(), + Name: peerID, + DNSLabel: peerID, + UserID: "testuser", + Status: &nbpeer.PeerStatus{Connected: false, LastSeen: time.Now()}, + SSHEnabled: false, + } + account.Peers[peerID] = peer + group, _ := account.GetGroupAll() + group.Peers = append(group.Peers, peerID) + user := &types.User{ + Id: fmt.Sprintf("%s-user-%d", account.Id, n), + AccountID: account.Id, + } + account.Users[user.Id] = user + route := &nbroute.Route{ + ID: nbroute.ID(fmt.Sprintf("network-id-%d", n)), + Description: "base route", + NetID: nbroute.NetID(fmt.Sprintf("network-id-%d", n)), + Network: netip.MustParsePrefix(netIP.String() + "/24"), + NetworkType: nbroute.IPv4Network, + Metric: 9999, + Masquerade: false, + Enabled: true, + Groups: []string{groupALL.ID}, + } + account.Routes[route.ID] = route + + group = &types.Group{ + ID: fmt.Sprintf("group-id-%d", n), + AccountID: account.Id, + Name: fmt.Sprintf("group-id-%d", n), + Issued: "api", + Peers: nil, + } + account.Groups[group.ID] = group + + nameserver := &nbdns.NameServerGroup{ + ID: fmt.Sprintf("nameserver-id-%d", n), + AccountID: account.Id, + Name: fmt.Sprintf("nameserver-id-%d", n), + Description: "", + NameServers: []nbdns.NameServer{{IP: netip.MustParseAddr(netIP.String()), NSType: nbdns.UDPNameServerType}}, + Groups: []string{group.ID}, + Primary: false, + Domains: nil, + Enabled: false, + SearchDomainsEnabled: false, + } + account.NameServerGroups[nameserver.ID] = nameserver + + setupKey, _ := types.GenerateDefaultSetupKey() + _, exists := account.SetupKeys[setupKey.Key] + if exists { + t.Errorf("setup key already exists") + } + account.SetupKeys[setupKey.Key] = setupKey + } + + err = store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + if len(store.GetAllAccounts(context.Background())) != 1 { + t.Errorf("expecting 1 Accounts to be stored after SaveAccount()") + } + + a, err := store.GetAccount(context.Background(), account.Id) + if a == nil { + t.Errorf("expecting Account to be stored after SaveAccount(): %v", err) + } + + if a != nil && len(a.Policies) != 1 { + t.Errorf("expecting Account to have one policy stored after SaveAccount(), got %d", len(a.Policies)) + } + + if a != nil && len(a.Policies[0].Rules) != 1 { + t.Errorf("expecting Account to have one policy rule stored after SaveAccount(), got %d", len(a.Policies[0].Rules)) + return + } + + if a != nil && len(a.Peers) != numPerAccount { + t.Errorf("expecting Account to have %d peers stored after SaveAccount(), got %d", + numPerAccount, len(a.Peers)) + return + } + + if a != nil && len(a.Users) != numPerAccount+1 { + t.Errorf("expecting Account to have %d users stored after SaveAccount(), got %d", + numPerAccount+1, len(a.Users)) + return + } + + if a != nil && len(a.Routes) != numPerAccount { + t.Errorf("expecting Account to have %d routes stored after SaveAccount(), got %d", + numPerAccount, len(a.Routes)) + return + } + + if a != nil && len(a.NameServerGroups) != numPerAccount { + t.Errorf("expecting Account to have %d NameServerGroups stored after SaveAccount(), got %d", + numPerAccount, len(a.NameServerGroups)) + return + } + + if a != nil && len(a.NameServerGroups) != numPerAccount { + t.Errorf("expecting Account to have %d NameServerGroups stored after SaveAccount(), got %d", + numPerAccount, len(a.NameServerGroups)) + return + } + + if a != nil && len(a.SetupKeys) != numPerAccount+1 { + t.Errorf("expecting Account to have %d SetupKeys stored after SaveAccount(), got %d", + numPerAccount+1, len(a.SetupKeys)) + return + } +} + +// sequentialIPv4 returns a unique IPv4 address for the given index, avoiding +// the random collisions that would otherwise violate the unique (account_id, ip) +// index when generating a large number of peers. +func sequentialIPv4(n int) net.IP { + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, 0x0A000000+uint32(n)) + return net.IP(b) +} + +func Test_SaveAccount(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("The SQLite store is not properly supported by Windows yet") + } + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + account := newAccountWithId(context.Background(), "account_id", "testuser", "") + setupKey, _ := types.GenerateDefaultSetupKey() + account.SetupKeys[setupKey.Key] = setupKey + account.Peers["testpeer"] = &nbpeer.Peer{ + Key: "peerkey", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::1"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + + err := store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + account2 := newAccountWithId(context.Background(), "account_id2", "testuser2", "") + setupKey, _ = types.GenerateDefaultSetupKey() + account2.SetupKeys[setupKey.Key] = setupKey + account2.Peers["testpeer2"] = &nbpeer.Peer{ + Key: "peerkey2", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 2}), + IPv6: netip.MustParseAddr("fd00::2"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name 2", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + + err = store.SaveAccount(context.Background(), account2) + require.NoError(t, err) + + if len(store.GetAllAccounts(context.Background())) != 2 { + t.Errorf("expecting 2 Accounts to be stored after SaveAccount()") + } + + a, err := store.GetAccount(context.Background(), account.Id) + if a == nil { + t.Errorf("expecting Account to be stored after SaveAccount(): %v", err) + } + + if a != nil && len(a.Policies) != 1 { + t.Errorf("expecting Account to have one policy stored after SaveAccount(), got %d", len(a.Policies)) + } + + if a != nil && len(a.Policies[0].Rules) != 1 { + t.Errorf("expecting Account to have one policy rule stored after SaveAccount(), got %d", len(a.Policies[0].Rules)) + return + } + + if a, err := store.GetAccountByPeerPubKey(context.Background(), "peerkey"); a == nil { + t.Errorf("expecting PeerKeyID2AccountID index updated after SaveAccount(): %v", err) + } + + if a, err := store.GetAccountByUser(context.Background(), "testuser"); a == nil { + t.Errorf("expecting UserID2AccountID index updated after SaveAccount(): %v", err) + } + + if a, err := store.GetAccountByPeerID(context.Background(), "testpeer"); a == nil { + t.Errorf("expecting PeerID2AccountID index updated after SaveAccount(): %v", err) + } + + if a, err := store.GetAccountBySetupKey(context.Background(), setupKey.Key); a == nil { + t.Errorf("expecting SetupKeyID2AccountID index updated after SaveAccount(): %v", err) + } + }) +} + +func Test_AccountSettings_SaveAndRetrieve(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("The SQLite store is not properly supported by Windows yet") + } + + populateFields := testing_helpers.NewPopulateFields().WithCustomFieldSetter( + reflect.PointerTo(reflect.TypeOf(types.ExtraSettings{})), func(this *testing_helpers.PopulateFields, field reflect.Value) (int, error) { + es := types.ExtraSettings{} + reflectedEs := reflect.ValueOf(&es).Elem() + n, err := this.PopulateAll(reflectedEs) + if err != nil { + return n, err + } + field.Set(reflectedEs.Addr()) + return n, nil + }).WithCustomFieldSetter( + reflect.PointerTo(reflect.TypeOf(types.DashboardFeatures{})), func(this *testing_helpers.PopulateFields, field reflect.Value) (int, error) { + t := true + df := types.DashboardFeatures{AgentNetwork: &t} + reflectedDf := reflect.ValueOf(&df).Elem() + field.Set(reflectedDf.Addr()) + return 1, nil + }).WithSkippedTag("gorm", "-") + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + account := newAccountWithId(context.Background(), "account_id", "testuser", "") + setupKey, _ := types.GenerateDefaultSetupKey() + account.SetupKeys[setupKey.Key] = setupKey + + settings := types.Settings{} + numOfExportedFields, err := populateFields.PopulateAll(reflect.ValueOf(&settings).Elem()) + assert.NoError(t, err) + assert.Equal(t, 27, numOfExportedFields) + account.Settings = &settings + + err = store.SaveAccount(context.Background(), account) + assert.NoError(t, err) + + accountFromDb, err := store.GetAccount(context.Background(), account.Id) + assert.NoError(t, err) + assert.NotNil(t, accountFromDb) + assert.NotNil(t, accountFromDb.Settings) + + assert.True(t, reflect.DeepEqual(&settings, accountFromDb.Settings), "created settings and settings retrieved from the db should match") + }) +} + +func TestSqlite_DeleteAccount(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("The SQLite store is not properly supported by Windows yet") + } + + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + testUserID := "testuser" + user := types.NewAdminUser(testUserID) + user.PATs = map[string]*types.PersonalAccessToken{"testtoken": { + ID: "testtoken", + Name: "test token", + }} + + account := newAccountWithId(context.Background(), "account_id", testUserID, "") + setupKey, _ := types.GenerateDefaultSetupKey() + account.SetupKeys[setupKey.Key] = setupKey + account.Peers["testpeer"] = &nbpeer.Peer{ + Key: "peerkey", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::1"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + account.Users[testUserID] = user + account.Networks = []*networkTypes.Network{ + { + ID: "network_id", + AccountID: account.Id, + Name: "network name", + Description: "network description", + }, + } + account.NetworkRouters = []*routerTypes.NetworkRouter{ + { + ID: "router_id", + NetworkID: account.Networks[0].ID, + AccountID: account.Id, + PeerGroups: []string{"group_id"}, + Masquerade: true, + Metric: 1, + }, + } + account.NetworkResources = []*resourceTypes.NetworkResource{ + { + ID: "resource_id", + NetworkID: account.Networks[0].ID, + AccountID: account.Id, + Name: "Name", + Description: "Description", + Type: "Domain", + Address: "example.com", + }, + } + + account.Services = []*rpservice.Service{ + { + ID: "service_id", + AccountID: account.Id, + Name: "test service", + Domain: "svc.example.com", + Enabled: true, + Targets: []*rpservice.Target{ + { + AccountID: account.Id, + ServiceID: "service_id", + Host: "localhost", + Port: 8080, + Protocol: "http", + Enabled: true, + }, + }, + }, + } + + account.Domains = []*proxydomain.Domain{ + { + ID: "domain_id", + Domain: "custom.example.com", + AccountID: account.Id, + Validated: true, + }, + } + + err = store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + if len(store.GetAllAccounts(context.Background())) != 1 { + t.Errorf("expecting 1 Accounts to be stored after SaveAccount()") + } + + o, err := store.GetAccountOnboarding(context.Background(), account.Id) + require.NoError(t, err) + require.Equal(t, o.AccountID, account.Id) + + err = store.CreateAgentNetworkSettings(context.Background(), &agentNetworkTypes.Settings{ + AccountID: account.Id, + Domain: "gw.example.com", + ProxyAddress: "gw.example.com", + }) + require.NoError(t, err) + + agentNetworkConfig := []any{ + &agentNetworkTypes.Provider{ID: "an_provider", AccountID: account.Id, APIKey: "sk-test"}, + &agentNetworkTypes.Policy{ID: "an_policy", AccountID: account.Id}, + &agentNetworkTypes.Guardrail{ID: "an_guardrail", AccountID: account.Id}, + &agentNetworkTypes.AccountBudgetRule{ID: "an_budget_rule", AccountID: account.Id}, + } + for _, row := range agentNetworkConfig { + require.NoError(t, store.(*SqlStore).db.Create(row).Error, "creating %T", row) + } + otherProvider := &agentNetworkTypes.Provider{ID: "other_provider", AccountID: "other_account"} + require.NoError(t, store.(*SqlStore).db.Create(otherProvider).Error) + + err = store.DeleteAccount(context.Background(), account) + require.NoError(t, err) + + _, err = store.GetAccountOnboarding(context.Background(), account.Id) + require.Error(t, err, "expecting error after removing DeleteAccount when getting onboarding") + + if len(store.GetAllAccounts(context.Background())) != 0 { + t.Errorf("expecting 0 Accounts to be stored after DeleteAccount()") + } + + _, err = store.GetAccountByPeerPubKey(context.Background(), "peerkey") + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer public key") + + _, err = store.GetAccountByUser(context.Background(), "testuser") + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by user") + + _, err = store.GetAccountByPeerID(context.Background(), "testpeer") + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer id") + + _, err = store.GetAccountBySetupKey(context.Background(), setupKey.Key) + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by setup key") + + _, err = store.GetAccount(context.Background(), account.Id) + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by id") + + for _, policy := range account.Policies { + var rules []*types.PolicyRule + err = store.(*SqlStore).db.Model(&types.PolicyRule{}).Find(&rules, "policy_id = ?", policy.ID).Error + require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for policy rules") + require.Len(t, rules, 0, "expecting no policy rules to be found after removing DeleteAccount") + + } + + for _, accountUser := range account.Users { + var pats []*types.PersonalAccessToken + err = store.(*SqlStore).db.Model(&types.PersonalAccessToken{}).Find(&pats, "user_id = ?", accountUser.Id).Error + require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for personal access token") + require.Len(t, pats, 0, "expecting no personal access token to be found after removing DeleteAccount") + + } + + for _, network := range account.Networks { + routers, err := store.GetNetworkRoutersByNetID(context.Background(), LockingStrengthNone, account.Id, network.ID) + require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for network routers") + require.Len(t, routers, 0, "expecting no network routers to be found after DeleteAccount") + + resources, err := store.GetNetworkResourcesByNetID(context.Background(), LockingStrengthNone, account.Id, network.ID) + require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for network resources") + require.Len(t, resources, 0, "expecting no network resources to be found after DeleteAccount") + } + + domains, err := store.ListCustomDomains(context.Background(), account.Id) + require.NoError(t, err, "expecting no error after DeleteAccount when searching for custom domains") + require.Len(t, domains, 0, "expecting no custom domains to be found after DeleteAccount") + + var services []*rpservice.Service + err = store.(*SqlStore).db.Model(&rpservice.Service{}).Find(&services, "account_id = ?", account.Id).Error + require.NoError(t, err, "expecting no error after DeleteAccount when searching for services") + require.Len(t, services, 0, "expecting no services to be found after DeleteAccount") + + var targets []*rpservice.Target + err = store.(*SqlStore).db.Model(&rpservice.Target{}).Find(&targets, "account_id = ?", account.Id).Error + require.NoError(t, err, "expecting no error after DeleteAccount when searching for service targets") + require.Len(t, targets, 0, "expecting no service targets to be found after DeleteAccount") + + _, err = store.GetAgentNetworkSettings(context.Background(), LockingStrengthNone, account.Id) + require.Error(t, err, "expecting agent network settings to be deleted with the account") + sErr, ok := status.FromError(err) + require.True(t, ok, "expecting a status error when getting agent network settings, got %v", err) + require.Equal(t, status.NotFound, sErr.Type(), "expecting agent network settings to be deleted with the account") + + // The domain is globally unique, so a leftover row would keep it from another account. + err = store.CreateAgentNetworkSettings(context.Background(), &agentNetworkTypes.Settings{ + AccountID: "other_account", + Domain: "gw.example.com", + ProxyAddress: "gw.example.com", + }) + require.NoError(t, err, "expecting the deleted account's gateway domain to be free for another account") + + for _, row := range agentNetworkConfig { + var count int64 + err = store.(*SqlStore).db.Model(row).Where("account_id = ?", account.Id).Count(&count).Error + require.NoError(t, err, "counting %T rows after DeleteAccount", row) + assert.Zero(t, count, "expecting no %T rows to be found after DeleteAccount", row) + } + + var otherProviders int64 + err = store.(*SqlStore).db.Model(&agentNetworkTypes.Provider{}).Where("account_id = ?", "other_account").Count(&otherProviders).Error + require.NoError(t, err) + assert.Equal(t, int64(1), otherProviders, "expecting another account's agent network provider to survive DeleteAccount") +} + +func Test_GetAccount(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("The SQLite store is not properly supported by Windows yet") + } + + runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { + id := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + account, err := store.GetAccount(context.Background(), id) + require.NoError(t, err) + require.Equal(t, id, account.Id, "account id should match") + require.Equal(t, false, account.Onboarding.OnboardingFlowPending) + + id = "9439-34653001fc3b-bf1c8084-ba50-4ce7" + + account, err = store.GetAccount(context.Background(), id) + require.NoError(t, err) + require.Equal(t, id, account.Id, "account id should match") + require.Equal(t, true, account.Onboarding.OnboardingFlowPending) + + _, err = store.GetAccount(context.Background(), "non-existing-account") + assert.Error(t, err) + parsedErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") + + }) +} + +func Test_TestGetAccountByPrivateDomain(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("The SQLite store is not properly supported by Windows yet") + } + + runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { + existingDomain := "test.com" + + account, err := store.GetAccountByPrivateDomain(context.Background(), existingDomain) + require.NoError(t, err, "should found account") + require.Equal(t, existingDomain, account.Domain, "domains should match") + + _, err = store.GetAccountByPrivateDomain(context.Background(), "missing-domain.com") + require.Error(t, err, "should return error on domain lookup") + parsedErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") + }) +} + +func TestPostgresql_SaveAccount(t *testing.T) { + if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { + t.Skip("skip CI tests on darwin and windows") + } + + t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + account := newAccountWithId(context.Background(), "account_id", "testuser", "") + setupKey, _ := types.GenerateDefaultSetupKey() + account.SetupKeys[setupKey.Key] = setupKey + account.Peers["testpeer"] = &nbpeer.Peer{ + Key: "peerkey", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::1"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + + err = store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + account2 := newAccountWithId(context.Background(), "account_id2", "testuser2", "") + setupKey, _ = types.GenerateDefaultSetupKey() + account2.SetupKeys[setupKey.Key] = setupKey + account2.Peers["testpeer2"] = &nbpeer.Peer{ + Key: "peerkey2", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 2}), + IPv6: netip.MustParseAddr("fd00::2"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name 2", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + + err = store.SaveAccount(context.Background(), account2) + require.NoError(t, err) + + if len(store.GetAllAccounts(context.Background())) != 2 { + t.Errorf("expecting 2 Accounts to be stored after SaveAccount()") + } + + a, err := store.GetAccount(context.Background(), account.Id) + if a == nil { + t.Errorf("expecting Account to be stored after SaveAccount(): %v", err) + } + + if a != nil && len(a.Policies) != 1 { + t.Errorf("expecting Account to have one policy stored after SaveAccount(), got %d", len(a.Policies)) + } + + if a != nil && len(a.Policies[0].Rules) != 1 { + t.Errorf("expecting Account to have one policy rule stored after SaveAccount(), got %d", len(a.Policies[0].Rules)) + return + } + + if a, err := store.GetAccountByPeerPubKey(context.Background(), "peerkey"); a == nil { + t.Errorf("expecting PeerKeyID2AccountID index updated after SaveAccount(): %v", err) + } + + if a, err := store.GetAccountByUser(context.Background(), "testuser"); a == nil { + t.Errorf("expecting UserID2AccountID index updated after SaveAccount(): %v", err) + } + + if a, err := store.GetAccountByPeerID(context.Background(), "testpeer"); a == nil { + t.Errorf("expecting PeerID2AccountID index updated after SaveAccount(): %v", err) + } + + if a, err := store.GetAccountBySetupKey(context.Background(), setupKey.Key); a == nil { + t.Errorf("expecting SetupKeyID2AccountID index updated after SaveAccount(): %v", err) + } +} + +func TestPostgresql_DeleteAccount(t *testing.T) { + if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { + t.Skip("skip CI tests on darwin and windows") + } + + t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + testUserID := "testuser" + user := types.NewAdminUser(testUserID) + user.PATs = map[string]*types.PersonalAccessToken{"testtoken": { + ID: "testtoken", + Name: "test token", + }} + + account := newAccountWithId(context.Background(), "account_id", testUserID, "") + setupKey, _ := types.GenerateDefaultSetupKey() + account.SetupKeys[setupKey.Key] = setupKey + account.Peers["testpeer"] = &nbpeer.Peer{ + Key: "peerkey", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::1"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + account.Users[testUserID] = user + + err = store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + if len(store.GetAllAccounts(context.Background())) != 1 { + t.Errorf("expecting 1 Accounts to be stored after SaveAccount()") + } + + err = store.DeleteAccount(context.Background(), account) + require.NoError(t, err) + + if len(store.GetAllAccounts(context.Background())) != 0 { + t.Errorf("expecting 0 Accounts to be stored after DeleteAccount()") + } + + _, err = store.GetAccountByPeerPubKey(context.Background(), "peerkey") + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer public key") + + _, err = store.GetAccountByUser(context.Background(), "testuser") + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by user") + + _, err = store.GetAccountByPeerID(context.Background(), "testpeer") + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer id") + + _, err = store.GetAccountBySetupKey(context.Background(), setupKey.Key) + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by setup key") + + _, err = store.GetAccount(context.Background(), account.Id) + require.Error(t, err, "expecting error after removing DeleteAccount when getting account by id") + + for _, policy := range account.Policies { + var rules []*types.PolicyRule + err = store.(*SqlStore).db.Model(&types.PolicyRule{}).Find(&rules, "policy_id = ?", policy.ID).Error + require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for policy rules") + require.Len(t, rules, 0, "expecting no policy rules to be found after removing DeleteAccount") + + } + + for _, accountUser := range account.Users { + var pats []*types.PersonalAccessToken + err = store.(*SqlStore).db.Model(&types.PersonalAccessToken{}).Find(&pats, "user_id = ?", accountUser.Id).Error + require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for personal access token") + require.Len(t, pats, 0, "expecting no personal access token to be found after removing DeleteAccount") + + } + +} + +func TestPostgresql_TestGetAccountByPrivateDomain(t *testing.T) { + if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { + t.Skip("skip CI tests on darwin and windows") + } + + t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + existingDomain := "test.com" + + account, err := store.GetAccountByPrivateDomain(context.Background(), existingDomain) + require.NoError(t, err, "should found account") + require.Equal(t, existingDomain, account.Domain, "domains should match") + + _, err = store.GetAccountByPrivateDomain(context.Background(), "missing-domain.com") + require.Error(t, err, "should return error on domain lookup") +} + +func TestSqlite_GetAccountNetwork(t *testing.T) { + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + if err != nil { + t.Fatal(err) + } + + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + _, err = store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + network, err := store.GetAccountNetwork(context.Background(), LockingStrengthNone, existingAccountID) + require.NoError(t, err) + ip := net.IP{100, 64, 0, 0}.To16() + assert.Equal(t, ip, network.Net.IP) + assert.Equal(t, net.IPMask{255, 255, 0, 0}, network.Net.Mask) + assert.Equal(t, "", network.Dns) + assert.Equal(t, "af1c8024-ha40-4ce2-9418-34653101fc3c", network.Identifier) + assert.Equal(t, uint64(0), network.Serial) +} + +func TestSqlStore_SaveAccountPersistsAgentNetworkOnly(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + account, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.False(t, account.Settings.AgentNetworkOnly, "setting should default to false") + + account.Settings.AgentNetworkOnly = true + require.NoError(t, store.SaveAccount(context.Background(), account)) + + reloaded, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.True(t, reloaded.Settings.AgentNetworkOnly, "setting should survive a save/load round-trip") + + reloaded.Settings.AgentNetworkOnly = false + require.NoError(t, store.SaveAccount(context.Background(), reloaded)) + + disabled, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.False(t, disabled.Settings.AgentNetworkOnly, "disabling should persist") +} + +func TestSqlStore_SaveAccountPersistsDashboardFeatures(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + account, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.Nil(t, account.Settings.DashboardFeatures, "dashboard features should default to unset") + + agentNetwork := true + account.Settings.DashboardFeatures = &types.DashboardFeatures{AgentNetwork: &agentNetwork} + require.NoError(t, store.SaveAccount(context.Background(), account)) + + reloaded, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.NotNil(t, reloaded.Settings.DashboardFeatures, "dashboard features should survive a save/load round-trip") + require.NotNil(t, reloaded.Settings.DashboardFeatures.AgentNetwork, "agent network flag should be set") + require.True(t, *reloaded.Settings.DashboardFeatures.AgentNetwork, "agent network flag should persist as true") + + disabled := false + reloaded.Settings.DashboardFeatures = &types.DashboardFeatures{AgentNetwork: &disabled} + require.NoError(t, store.SaveAccount(context.Background(), reloaded)) + + reloadedDisabled, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.NotNil(t, reloadedDisabled.Settings.DashboardFeatures.AgentNetwork, "agent network flag should remain set") + require.False(t, *reloadedDisabled.Settings.DashboardFeatures.AgentNetwork, "explicit false should persist") +} + +func TestSqlStore_UpdateAccountDomainAttributes(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + if err != nil { + t.Fatal(err) + } + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + t.Run("Should update attributes with public domain", func(t *testing.T) { + require.NoError(t, err) + domain := "example.com" + category := "public" + IsDomainPrimaryAccount := false + err = store.UpdateAccountDomainAttributes(context.Background(), accountID, domain, category, IsDomainPrimaryAccount) + require.NoError(t, err) + account, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.Equal(t, domain, account.Domain) + require.Equal(t, category, account.DomainCategory) + require.Equal(t, IsDomainPrimaryAccount, account.IsDomainPrimaryAccount) + }) + + t.Run("Should update attributes with private domain", func(t *testing.T) { + require.NoError(t, err) + domain := "test.com" + category := "private" + IsDomainPrimaryAccount := true + err = store.UpdateAccountDomainAttributes(context.Background(), accountID, domain, category, IsDomainPrimaryAccount) + require.NoError(t, err) + account, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + require.Equal(t, domain, account.Domain) + require.Equal(t, category, account.DomainCategory) + require.Equal(t, IsDomainPrimaryAccount, account.IsDomainPrimaryAccount) + }) + + t.Run("Should fail when account does not exist", func(t *testing.T) { + require.NoError(t, err) + domain := "test.com" + category := "private" + IsDomainPrimaryAccount := true + err = store.UpdateAccountDomainAttributes(context.Background(), "non-existing-account-id", domain, category, IsDomainPrimaryAccount) + require.Error(t, err) + }) + +} + +func TestSqlStore_GetDNSSettings(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectError bool + }{ + { + name: "retrieve existing account dns settings", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectError: false, + }, + { + name: "retrieve non-existing account dns settings", + accountID: "non-existing", + expectError: true, + }, + { + name: "retrieve dns settings with empty account ID", + accountID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dnsSettings, err := store.GetAccountDNSSettings(context.Background(), LockingStrengthNone, tt.accountID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, dnsSettings) + } else { + require.NoError(t, err) + require.NotNil(t, dnsSettings) + } + }) + } +} + +func TestSqlStore_SaveDNSSettings(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + dnsSettings, err := store.GetAccountDNSSettings(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + + dnsSettings.DisabledManagementGroups = []string{"groupA", "groupB"} + err = store.SaveDNSSettings(context.Background(), accountID, dnsSettings) + require.NoError(t, err) + + saveDNSSettings, err := store.GetAccountDNSSettings(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Equal(t, saveDNSSettings, dnsSettings) +} + +func TestSqlStore_GetAccountCreatedBy(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectError bool + createdBy string + }{ + { + name: "existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectError: false, + createdBy: "edafee4e-63fb-11ec-90d6-0242ac120003", + }, + { + name: "non-existing account ID", + accountID: "nonexistent", + expectError: true, + }, + { + name: "empty account ID", + accountID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + createdBy, err := store.GetAccountCreatedBy(context.Background(), LockingStrengthNone, tt.accountID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Empty(t, createdBy) + } else { + require.NoError(t, err) + require.NotNil(t, createdBy) + require.Equal(t, tt.createdBy, createdBy) + } + }) + } + +} + +func TestSqlStore_GetAccountMeta(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + accountMeta, err := store.GetAccountMeta(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.NotNil(t, accountMeta) + require.Equal(t, accountID, accountMeta.AccountID) + require.Equal(t, "edafee4e-63fb-11ec-90d6-0242ac120003", accountMeta.CreatedBy) + require.Equal(t, "test.com", accountMeta.Domain) + require.Equal(t, "private", accountMeta.DomainCategory) + require.Equal(t, time.Date(2024, time.October, 2, 14, 1, 38, 210000000, time.UTC), accountMeta.CreatedAt.UTC()) +} + +func TestSqlStore_GetAnyAccountID(t *testing.T) { + t.Run("should return account ID when accounts exist", func(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID, err := store.GetAnyAccountID(context.Background()) + require.NoError(t, err) + assert.Equal(t, "bf1c8084-ba50-4ce7-9439-34653001fc3b", accountID) + }) + + t.Run("should return error when no accounts exist", func(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID, err := store.GetAnyAccountID(context.Background()) + require.Error(t, err) + sErr, ok := status.FromError(err) + assert.True(t, ok) + assert.Equal(t, sErr.Type(), status.NotFound) + assert.Empty(t, accountID) + }) +} diff --git a/management/server/store/sql_store_agent_network_access_log.go b/management/server/store/sql_store_agent_network_access_log.go new file mode 100644 index 000000000..479dfa45e --- /dev/null +++ b/management/server/store/sql_store_agent_network_access_log.go @@ -0,0 +1,285 @@ +package store + +import ( + "context" + "time" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// CreateAgentNetworkAccessLog persists a flattened agent-network access-log +// entry together with its authorising-group child rows in a single +// transaction. +func (s *SqlStore) CreateAgentNetworkAccessLog(ctx context.Context, entry *agentNetworkTypes.AgentNetworkAccessLog, groups []agentNetworkTypes.AgentNetworkAccessLogGroup) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { + // Idempotent on the log id / (log_id, group_id) so a proxy resend of the + // same entry can't fail the request. + if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(entry).Error; err != nil { + return err + } + if len(groups) > 0 { + if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&groups).Error; err != nil { + return err + } + } + return nil + }) + if err != nil { + log.WithContext(ctx).WithFields(log.Fields{ + "account_id": entry.AccountID, + "service_id": entry.ServiceID, + "model": entry.Model, + }).Errorf("failed to create agent-network access log entry in store: %v", err) + return status.Errorf(status.Internal, "failed to create agent-network access log entry in store") + } + return nil +} + +// DeleteOldAgentNetworkAccessLogs deletes an account's access-log rows (and +// their authorising-group child rows) older than the cutoff. Usage records are +// untouched — they are the long-term aggregate. Returns the number of log rows +// deleted. +func (s *SqlStore) DeleteOldAgentNetworkAccessLogs(ctx context.Context, accountID string, olderThan time.Time) (int64, error) { + var deleted int64 + err := s.transaction(ctx, func(tx *gorm.DB) error { + // Remove group child rows for the soon-to-be-deleted logs first. + if err := tx.Exec( + "DELETE FROM agent_network_access_log_group WHERE account_id = ? AND log_id IN (SELECT id FROM agent_network_access_log WHERE account_id = ? AND timestamp < ?)", + accountID, accountID, olderThan, + ).Error; err != nil { + return err + } + res := tx.Where("account_id = ? AND timestamp < ?", accountID, olderThan). + Delete(&agentNetworkTypes.AgentNetworkAccessLog{}) + if res.Error != nil { + return res.Error + } + deleted = res.RowsAffected + return nil + }) + if err != nil { + log.WithContext(ctx).Errorf("failed to delete old agent-network access logs for account %s: %v", accountID, err) + return 0, status.Errorf(status.Internal, "failed to delete old agent-network access logs") + } + return deleted, nil +} + +// GetDeletedAccountIDsWithAgentNetworkAccessLogs returns the IDs of accounts that no +// longer exist but still have access-log rows. The retention sweep is driven by settings +// rows, which are deleted with the account, so it uses this to find logs it would +// otherwise never expire. +func (s *SqlStore) GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx context.Context) ([]string, error) { + var accountIDs []string + err := s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}). + Distinct("account_id"). + Where("NOT EXISTS (SELECT 1 FROM accounts WHERE accounts.id = agent_network_access_log.account_id)"). + Pluck("account_id", &accountIDs).Error + if err != nil { + log.WithContext(ctx).Errorf("failed to get deleted accounts with agent-network access logs: %v", err) + return nil, status.Errorf(status.Internal, "failed to get deleted accounts with agent-network access logs") + } + return accountIDs, nil +} + +// GetAgentNetworkAccessLogs retrieves flattened agent-network access logs for +// an account with server-side pagination, filtering and sorting. Authorising +// group ids are hydrated from the group child table for the returned page. +func (s *SqlStore) GetAgentNetworkAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLog, int64, error) { + var logs []*agentNetworkTypes.AgentNetworkAccessLog + var totalCount int64 + + countQuery := s.applyAgentNetworkAccessLogFilters( + s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}).Where(accountIDCondition, accountID), + filter, + ) + if err := countQuery.Count(&totalCount).Error; err != nil { + log.WithContext(ctx).Errorf("failed to count agent-network access logs: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to count agent-network access logs") + } + + query := s.applyAgentNetworkAccessLogFilters( + s.db.Where(accountIDCondition, accountID), + filter, + ). + Order(filter.GetSortColumn() + " " + filter.GetSortOrder()). + Limit(filter.GetLimit()). + Offset(filter.GetOffset()) + + if lockStrength != LockingStrengthNone { + query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + if err := query.Find(&logs).Error; err != nil { + log.WithContext(ctx).Errorf("failed to get agent-network access logs from store: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to get agent-network access logs from store") + } + + if err := s.hydrateAgentNetworkAccessLogGroups(ctx, accountID, logs); err != nil { + return nil, 0, err + } + + return logs, totalCount, nil +} + +// applyAgentNetworkAccessLogFilters applies the filter conditions to a query. +func (s *SqlStore) applyAgentNetworkAccessLogFilters(query *gorm.DB, filter agentNetworkTypes.AgentNetworkAccessLogFilter) *gorm.DB { + if filter.Search != nil { + p := "%" + *filter.Search + "%" + query = query.Where( + "id LIKE ? OR host LIKE ? OR path LIKE ? OR model LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)", + p, p, p, p, p, p, + ) + } + if filter.UserID != nil { + query = query.Where("user_id = ?", *filter.UserID) + } + if filter.SessionID != nil { + query = query.Where("session_id = ?", *filter.SessionID) + } + if filter.Decision != nil { + query = query.Where("decision = ?", *filter.Decision) + } + if filter.PathPrefix != nil { + query = query.Where("path LIKE ?", *filter.PathPrefix+"%") + } + if len(filter.ProviderIDs) > 0 { + query = query.Where("resolved_provider_id IN ?", filter.ProviderIDs) + } + if len(filter.Models) > 0 { + query = query.Where("model IN ?", filter.Models) + } + if len(filter.GroupIDs) > 0 { + query = query.Where( + "id IN (SELECT log_id FROM agent_network_access_log_group WHERE group_id IN ?)", + filter.GroupIDs, + ) + } + if filter.StartDate != nil { + query = query.Where("timestamp >= ?", *filter.StartDate) + } + if filter.EndDate != nil { + query = query.Where("timestamp <= ?", *filter.EndDate) + } + return query +} + +// hydrateAgentNetworkAccessLogGroups loads the authorising group ids for the +// given page of entries and assigns them onto each entry's GroupIDs field. +func (s *SqlStore) hydrateAgentNetworkAccessLogGroups(ctx context.Context, accountID string, logs []*agentNetworkTypes.AgentNetworkAccessLog) error { + if len(logs) == 0 { + return nil + } + + ids := make([]string, 0, len(logs)) + for _, l := range logs { + ids = append(ids, l.ID) + } + + var rows []agentNetworkTypes.AgentNetworkAccessLogGroup + if err := s.db. + Where(accountIDCondition, accountID). + Where("log_id IN ?", ids). + Find(&rows).Error; err != nil { + log.WithContext(ctx).Errorf("failed to hydrate agent-network access log groups: %v", err) + return status.Errorf(status.Internal, "failed to hydrate agent-network access log groups") + } + + byLog := make(map[string][]string, len(logs)) + for _, r := range rows { + byLog[r.LogID] = append(byLog[r.LogID], r.GroupID) + } + for _, l := range logs { + l.GroupIDs = byLog[l.ID] + } + return nil +} + +// agentNetworkSessionKeyExpr is the SQL group key for session-grouped access +// logs: the row's session id, or — when the client sent none — the row id, so +// session-less requests each form their own singleton group. COALESCE/NULLIF +// are standard SQL, so this stays portable across SQLite and Postgres. +const agentNetworkSessionKeyExpr = "COALESCE(NULLIF(session_id, ''), id)" + +// GetAgentNetworkAccessLogSessions retrieves agent-network access logs grouped +// by session, with server-side pagination, filtering and sorting at the session +// level. It paginates over the distinct session keys (ordered by the requested +// session-level aggregate), fetches every entry for the page's sessions, and +// folds them into per-session summaries. The returned count is the number of +// matching sessions. Filters apply to the entries, so a session's summary +// reflects only its filter-matching requests. +func (s *SqlStore) GetAgentNetworkAccessLogSessions(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLogSession, int64, error) { + // Count distinct sessions via a grouped subquery — portable and avoids + // relying on COUNT(DISTINCT ) quoting quirks. + sessionsSubquery := s.applyAgentNetworkAccessLogFilters( + s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}).Where(accountIDCondition, accountID), + filter, + ). + Select(agentNetworkSessionKeyExpr + " AS session_key"). + Group(agentNetworkSessionKeyExpr) + + var totalCount int64 + if err := s.db.Table("(?) AS sessions", sessionsSubquery).Count(&totalCount).Error; err != nil { + log.WithContext(ctx).Errorf("failed to count agent-network access-log sessions: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to count agent-network access-log sessions") + } + + // The page of session keys, ordered by the session-level aggregate. The + // session-key tiebreaker keeps pagination deterministic when the primary + // aggregate ties. + type sessionKeyRow struct { + SessionKey string + } + var keyRows []sessionKeyRow + keyQuery := s.applyAgentNetworkAccessLogFilters( + s.db.Model(&agentNetworkTypes.AgentNetworkAccessLog{}).Where(accountIDCondition, accountID), + filter, + ). + Select(agentNetworkSessionKeyExpr + " AS session_key"). + Group(agentNetworkSessionKeyExpr). + Order(filter.GetSessionSortExpr() + " " + filter.GetSortOrder()). + Order("session_key ASC"). + Limit(filter.GetLimit()). + Offset(filter.GetOffset()) + if err := keyQuery.Scan(&keyRows).Error; err != nil { + log.WithContext(ctx).Errorf("failed to list agent-network access-log session keys: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to list agent-network access-log session keys") + } + if len(keyRows) == 0 { + return nil, totalCount, nil + } + + keys := make([]string, 0, len(keyRows)) + for _, r := range keyRows { + keys = append(keys, r.SessionKey) + } + + // All entries for the page's sessions, contiguous per session and oldest + // first within each — the fold relies on that ordering. + var entries []*agentNetworkTypes.AgentNetworkAccessLog + entriesQuery := s.applyAgentNetworkAccessLogFilters( + s.db.Where(accountIDCondition, accountID), + filter, + ). + Where(agentNetworkSessionKeyExpr+" IN ?", keys). + Order(agentNetworkSessionKeyExpr + ", timestamp ASC") + + if lockStrength != LockingStrengthNone { + entriesQuery = entriesQuery.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + if err := entriesQuery.Find(&entries).Error; err != nil { + log.WithContext(ctx).Errorf("failed to get agent-network access-log session entries: %v", err) + return nil, 0, status.Errorf(status.Internal, "failed to get agent-network access-log session entries") + } + + if err := s.hydrateAgentNetworkAccessLogGroups(ctx, accountID, entries); err != nil { + return nil, 0, err + } + + return agentNetworkTypes.FoldAccessLogSessions(keys, entries), totalCount, nil +} diff --git a/management/server/store/sql_store_agent_network_usage.go b/management/server/store/sql_store_agent_network_usage.go new file mode 100644 index 000000000..dc557d2f4 --- /dev/null +++ b/management/server/store/sql_store_agent_network_usage.go @@ -0,0 +1,91 @@ +package store + +import ( + "context" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// CreateAgentNetworkUsage persists a stripped agent-network usage record +// together with its authorising-group child rows in a single transaction. +func (s *SqlStore) CreateAgentNetworkUsage(ctx context.Context, usage *agentNetworkTypes.AgentNetworkUsage, groups []agentNetworkTypes.AgentNetworkUsageGroup) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { + // Idempotent on the usage id / (usage_id, group_id) so a proxy resend of + // the same entry can't fail the request. + if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(usage).Error; err != nil { + return err + } + if len(groups) > 0 { + if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&groups).Error; err != nil { + return err + } + } + return nil + }) + if err != nil { + log.WithContext(ctx).WithFields(log.Fields{ + "account_id": usage.AccountID, + "model": usage.Model, + }).Errorf("failed to create agent-network usage record in store: %v", err) + return status.Errorf(status.Internal, "failed to create agent-network usage record in store") + } + return nil +} + +// GetAgentNetworkUsageRows returns the stripped usage rows for an account that +// match the filter (date / user / group / provider / model). Aggregation into +// time buckets happens in the manager so granularities stay engine-portable. +func (s *SqlStore) GetAgentNetworkUsageRows(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkUsage, error) { + var rows []*agentNetworkTypes.AgentNetworkUsage + + query := s.applyAgentNetworkUsageFilters( + s.db.Where(accountIDCondition, accountID), + filter, + ).Order("timestamp ASC") + + if lockStrength != LockingStrengthNone { + query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + if err := query.Find(&rows).Error; err != nil { + log.WithContext(ctx).Errorf("failed to get agent-network usage rows from store: %v", err) + return nil, status.Errorf(status.Internal, "failed to get agent-network usage rows from store") + } + return rows, nil +} + +// applyAgentNetworkUsageFilters applies the shared access-log filter's +// date/user/group/provider/model conditions to a usage-table query. Pagination, +// sort and free-text search are ignored — the overview is an aggregate. +func (s *SqlStore) applyAgentNetworkUsageFilters(query *gorm.DB, filter agentNetworkTypes.AgentNetworkAccessLogFilter) *gorm.DB { + if filter.UserID != nil { + query = query.Where("user_id = ?", *filter.UserID) + } + if filter.SessionID != nil { + query = query.Where("session_id = ?", *filter.SessionID) + } + if len(filter.ProviderIDs) > 0 { + query = query.Where("resolved_provider_id IN ?", filter.ProviderIDs) + } + if len(filter.Models) > 0 { + query = query.Where("model IN ?", filter.Models) + } + if len(filter.GroupIDs) > 0 { + query = query.Where( + "id IN (SELECT usage_id FROM agent_network_request_usage_group WHERE group_id IN ?)", + filter.GroupIDs, + ) + } + if filter.StartDate != nil { + query = query.Where("timestamp >= ?", *filter.StartDate) + } + if filter.EndDate != nil { + query = query.Where("timestamp <= ?", *filter.EndDate) + } + return query +} diff --git a/management/server/store/sql_store_agentnetwork.go b/management/server/store/sql_store_agentnetwork.go index 8a92f7147..bdb1f97d6 100644 --- a/management/server/store/sql_store_agentnetwork.go +++ b/management/server/store/sql_store_agentnetwork.go @@ -619,7 +619,7 @@ func (s *SqlStore) IncrementAgentNetworkConsumptionBatch( } const tbl = "agent_network_consumption" - err := s.db.Transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { for _, k := range keys { if k.DimID == "" || k.WindowSeconds <= 0 { return status.Errorf(status.InvalidArgument, "dim_id and window_seconds must be set") @@ -663,6 +663,21 @@ func (s *SqlStore) IncrementAgentNetworkConsumptionBatch( return nil } +// DeleteAgentNetworkConsumptionOfDeletedAccounts deletes every consumption counter whose +// account no longer exists and returns the number of rows deleted. Counters grow with +// traffic, so they are swept in the background instead of in the account-deletion +// transaction, and the sweep also catches counters a proxy writes after the deletion. +func (s *SqlStore) DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx context.Context) (int64, error) { + res := s.db. + Where("NOT EXISTS (SELECT 1 FROM accounts WHERE accounts.id = agent_network_consumption.account_id)"). + Delete(&agentNetworkTypes.Consumption{}) + if res.Error != nil { + log.WithContext(ctx).Errorf("failed to delete agent-network consumption of deleted accounts: %v", res.Error) + return 0, status.Errorf(status.Internal, "failed to delete agent-network consumption of deleted accounts") + } + return res.RowsAffected, nil +} + // ListAgentNetworkConsumption returns every consumption row recorded // for the account, ordered by window_start descending. Backs the // dashboard's basic counter view. diff --git a/management/server/store/sql_store_agentnetwork_accesslog_test.go b/management/server/store/sql_store_agentnetwork_accesslog_test.go index 8ba79a062..b8a3560e0 100644 --- a/management/server/store/sql_store_agentnetwork_accesslog_test.go +++ b/management/server/store/sql_store_agentnetwork_accesslog_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/types" ) // TestAgentNetworkUsage_RealStore_RoundTrip drives CreateAgentNetworkUsage and @@ -300,3 +301,37 @@ func TestDeleteOldAgentNetworkAccessLogs(t *testing.T) { require.NoError(t, err) require.Len(t, usage, 1, "usage record for the deleted log must survive") } + +// TestDeleteAgentNetworkConsumptionOfDeletedAccounts verifies that the sweep removes the +// consumption counters of accounts that no longer exist and leaves live accounts' counters, +// including those of a live account without a settings row. +func TestDeleteAgentNetworkConsumptionOfDeletedAccounts(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, s Store) { + ctx := context.Background() + const ( + liveAccountID = "acc-anet-consumption-live" + deletedAccountID = "acc-anet-consumption-deleted" + ) + require.NoError(t, s.SaveAccount(ctx, &types.Account{Id: liveAccountID})) + + windowStart := time.Now().UTC().Truncate(time.Hour) + for _, accountID := range []string{liveAccountID, deletedAccountID} { + for _, dimID := range []string{"user-1", "user-2"} { + require.NoError(t, s.IncrementAgentNetworkConsumption(ctx, accountID, + agentNetworkTypes.DimensionUser, dimID, 3600, windowStart, 10, 5, 0.01)) + } + } + + deleted, err := s.DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx) + require.NoError(t, err) + assert.Equal(t, int64(2), deleted, "both of the deleted account's counters should be removed") + + rows, err := s.ListAgentNetworkConsumption(ctx, LockingStrengthNone, deletedAccountID) + require.NoError(t, err) + assert.Empty(t, rows, "the deleted account should have no consumption counters left") + + rows, err = s.ListAgentNetworkConsumption(ctx, LockingStrengthNone, liveAccountID) + require.NoError(t, err) + assert.Len(t, rows, 2, "the live account's consumption counters should survive") + }) +} diff --git a/management/server/store/sql_store_custom_domain.go b/management/server/store/sql_store_custom_domain.go new file mode 100644 index 000000000..5b7cbf519 --- /dev/null +++ b/management/server/store/sql_store_custom_domain.go @@ -0,0 +1,178 @@ +package store + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/rs/xid" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" + "github.com/netbirdio/netbird/shared/management/status" +) + +// GetCustomDomainsCounts returns the total and validated custom domain counts. +func (s *SqlStore) GetCustomDomainsCounts(ctx context.Context) (int64, int64, error) { + var total, validated int64 + if err := s.db.Model(&domain.Domain{}).Count(&total).Error; err != nil { + return 0, 0, err + } + if err := s.db.Model(&domain.Domain{}).Where("validated = ?", true).Count(&validated).Error; err != nil { + return 0, 0, err + } + return total, validated, nil +} + +func (s *SqlStore) GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error) { + tx := s.db + + customDomain := &domain.Domain{} + result := tx.Take(&customDomain, accountAndIDQueryCondition, accountID, domainID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "custom domain %s not found", domainID) + } + + log.WithContext(ctx).Errorf("failed to get custom domain from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get custom domain from store") + } + + return customDomain, nil +} + +func (s *SqlStore) ListFreeDomains(ctx context.Context, accountID string) ([]string, error) { + return nil, nil +} + +func (s *SqlStore) ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) { + tx := s.db + + var domains []*domain.Domain + result := tx.Find(&domains, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get reverse proxy custom domains from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get reverse proxy custom domains from store") + } + + return domains, nil +} + +// GetCustomDomainByName returns the custom domain row holding the given name, +// regardless of which account owns it. +func (s *SqlStore) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) { + customDomain := &domain.Domain{} + result := s.db.Take(customDomain, "domain = ?", domainName) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "custom domain %s not found", domainName) + } + + log.WithContext(ctx).Errorf("failed to get custom domain by name from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get custom domain from store") + } + + return customDomain, nil +} + +func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) { + newDomain := &domain.Domain{ + ID: xid.New().String(), // Generate our own ID because gorm doesn't always configure the database to handle this for us. + Domain: domainName, + AccountID: accountID, + TargetCluster: targetCluster, + Type: domain.TypeCustom, + Validated: validated, + } + if !validated { + expiresAt := time.Now().UTC().Add(domain.ValidationTTL) + newDomain.ValidationExpiresAt = &expiresAt + } + result := s.db.Create(newDomain) + if result.Error != nil { + // The unique index is the last guard when two requests clear the + // manager's availability check at the same time. The one that loses the + // insert is a conflict, not an internal failure. + var count int64 + if err := s.db.Model(&domain.Domain{}).Where("domain = ?", domainName).Count(&count).Error; err == nil && count > 0 { + // The insert error is logged even on this path: the name being taken + // is what the caller has to act on, but if the insert also failed for + // an unrelated reason the operator still needs to see it. + log.WithContext(ctx).Warnf("create reverse proxy custom domain %s rejected, name already registered: %v", domainName, result.Error) + return nil, status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName) + } + + log.WithContext(ctx).Errorf("failed to create reverse proxy custom domain to store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to create reverse proxy custom domain to store") + } + + return newDomain, nil +} + +// UpdateCustomDomain completes validation only while the original registration is pending. +func (s *SqlStore) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) { + if !d.Validated { + return nil, status.Errorf(status.InvalidArgument, "custom domain update must complete validation") + } + result := s.db.WithContext(ctx).Model(&domain.Domain{}). + Where(accountAndIDQueryCondition, accountID, d.ID). + Where("domain = ? AND target_cluster = ?", d.Domain, d.TargetCluster). + Where("validated = ? AND validation_expires_at > ?", false, time.Now().UTC()). + Update("validated", true) + if result.Error != nil { + return nil, fmt.Errorf("validate custom domain in store: %w", result.Error) + } + if result.RowsAffected == 0 { + return nil, status.Errorf(status.PreconditionFailed, "custom domain registration is no longer pending validation") + } + + return d, nil +} + +// LockCustomDomains holds shared locks on registrations covering a service until commit. +func (s *SqlStore) LockCustomDomains(ctx context.Context, accountID string, serviceDomain nbdomain.Domain) ([]*domain.Domain, error) { + var names []string + for name := serviceDomain.PunycodeString(); name != ""; { + names = append(names, name) + _, name, _ = strings.Cut(name, ".") + } + + var domains []*domain.Domain + if err := s.db.WithContext(ctx).Clauses(clause.Locking{Strength: string(LockingStrengthShare)}). + Where(accountIDCondition, accountID).Where("domain IN ?", names). + Order("id").Find(&domains).Error; err != nil { + return nil, fmt.Errorf("lock custom domains: %w", err) + } + return domains, nil +} + +// DeleteCustomDomain removes a registration only when no service uses its namespace. +func (s *SqlStore) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error { + return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var d domain.Domain + // Service writes hold a shared lock on this row through commit, so neither + // operation can proceed against the other's outdated view of the domain. + if err := tx.Clauses(clause.Locking{Strength: string(LockingStrengthUpdate)}). + Take(&d, accountAndIDQueryCondition, accountID, domainID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return status.Errorf(status.NotFound, "custom domain not found") + } + return fmt.Errorf("lock custom domain for deletion: %w", err) + } + + result := tx.Where(accountAndIDQueryCondition, accountID, domainID). + Where("NOT EXISTS (?)", customDomainServices(tx, &d).Select("1")).Delete(&domain.Domain{}) + if result.Error != nil { + return fmt.Errorf("delete custom domain: %w", result.Error) + } + if result.RowsAffected == 0 { + return status.Errorf(status.PreconditionFailed, "custom domain has dependent services; delete or move them before deleting the domain") + } + return nil + }) +} diff --git a/management/server/store/sql_store_custom_domain_test.go b/management/server/store/sql_store_custom_domain_test.go new file mode 100644 index 000000000..a274c7271 --- /dev/null +++ b/management/server/store/sql_store_custom_domain_test.go @@ -0,0 +1,182 @@ +package store + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/types" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestLockCustomDomains_ConcurrentServices(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + if store.GetStoreEngine() == types.SqliteStoreEngine { + t.Skip("SQLite serializes transactions on one connection") + } + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + _, err := store.CreateCustomDomain(ctx, "owner", "one.example.com", "cluster", true) + require.NoError(t, err) + _, err = store.CreateCustomDomain(ctx, "owner", "two.example.com", "cluster", true) + require.NoError(t, err) + + locked := make(chan error, 1) + release := make(chan struct{}) + done := make(chan error, 1) + go func() { + done <- store.ExecuteInTransaction(ctx, func(tx Store) error { + _, err := tx.LockCustomDomains(ctx, "owner", "app.one.example.com") + locked <- err + if err != nil { + return err + } + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + }) + }() + var lockErr error + select { + case lockErr = <-locked: + case err := <-done: + t.Fatalf("transaction ended before locking: %v", err) + } + writeCtx, writeCancel := context.WithTimeout(ctx, 3*time.Second) + defer writeCancel() + var writeErr error + for _, name := range []nbdomain.Domain{"app.one.example.com", "app.two.example.com"} { + writeErr = store.ExecuteInTransaction(writeCtx, func(tx Store) error { + if _, err := tx.LockCustomDomains(writeCtx, "owner", name); err != nil { + return err + } + return tx.CreateService(writeCtx, &rpservice.Service{ + ID: name.PunycodeString(), AccountID: "owner", Domain: name.PunycodeString(), + }) + }) + if writeErr != nil { + break + } + } + close(release) + require.NoError(t, <-done) + require.NoError(t, lockErr) + require.NoError(t, writeErr, "domain authorization locks must allow concurrent service writes") + services, err := store.GetAccountServices(ctx, LockingStrengthNone, "owner") + require.NoError(t, err) + assert.Len(t, services, 2, "both services must commit while the first domain is locked") + }) +} + +func TestDeleteCustomDomain_ServiceDependencies(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + d, err := store.CreateCustomDomain(ctx, "owner", "example.com", "cluster", true) + require.NoError(t, err) + svc := &rpservice.Service{ID: "service", AccountID: "owner", Domain: "APP.EXAMPLE.COM."} + require.NoError(t, store.CreateService(ctx, svc)) + + err = store.DeleteCustomDomain(ctx, "other", d.ID) + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok, "cross-account deletion must return a typed error") + assert.Equal(t, status.NotFound, sErr.Type(), "cross-account deletion must not reveal dependencies") + + err = store.DeleteCustomDomain(ctx, "owner", d.ID) + require.Error(t, err) + sErr, ok = status.FromError(err) + require.True(t, ok, "dependent services must return a typed error") + assert.Equal(t, status.PreconditionFailed, sErr.Type(), "deletion must fail until services are removed") + stored, err := store.GetCustomDomain(ctx, "owner", d.ID) + require.NoError(t, err) + assert.True(t, stored.Validated, "rejected deletion must preserve validation") + + require.NoError(t, store.DeleteService(ctx, "owner", svc.ID)) + require.NoError(t, store.DeleteCustomDomain(ctx, "owner", d.ID)) + _, err = store.GetCustomDomain(ctx, "owner", d.ID) + require.Error(t, err, "the registration must be gone after successful deletion") + }) +} + +func TestDeleteCustomDomain_OtherAccountSubdomainIsNotADependency(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "other", "admin", ""))) + + // Registrations are unique by name, so a second account can hold a + // subdomain of the first account's registration and serve from it. + parent, err := store.CreateCustomDomain(ctx, "owner", "example.com", "cluster", true) + require.NoError(t, err) + _, err = store.CreateCustomDomain(ctx, "other", "team.example.com", "cluster", true) + require.NoError(t, err) + require.NoError(t, store.CreateService(ctx, &rpservice.Service{ + ID: "service", AccountID: "other", Domain: "app.team.example.com", + })) + + require.NoError(t, store.DeleteCustomDomain(ctx, "owner", parent.ID), + "another account's service must not hold the registration open") + }) +} + +func TestDeleteCustomDomain_ConcurrentServiceCreation(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + for i := range 10 { + d, err := store.CreateCustomDomain(ctx, "owner", fmt.Sprintf("app%d.example.com", i), "cluster", true) + require.NoError(t, err) + svc := &rpservice.Service{ID: fmt.Sprintf("service-%d", i), AccountID: "owner", Domain: "nested." + d.Domain} + start := make(chan struct{}) + created := make(chan error, 1) + deleted := make(chan error, 1) + go func() { + <-start + created <- store.ExecuteInTransaction(ctx, func(tx Store) error { + domains, err := tx.LockCustomDomains(ctx, "owner", nbdomain.Domain(svc.Domain)) + if err != nil { + return err + } + for _, candidate := range domains { + if candidate.ID == d.ID && candidate.Validated { + return tx.CreateService(ctx, svc) + } + } + return status.Errorf(status.PreconditionFailed, "registration was deleted") + }) + }() + go func() { + <-start + deleted <- store.DeleteCustomDomain(ctx, "owner", d.ID) + }() + close(start) + createErr, deleteErr := <-created, <-deleted + require.True(t, createErr == nil || deleteErr == nil, "one operation must succeed: create=%v, delete=%v", createErr, deleteErr) + if createErr == nil { + require.Error(t, deleteErr, "a committed service must block deletion") + stored, err := store.GetCustomDomain(ctx, "owner", d.ID) + require.NoError(t, err) + assert.True(t, stored.Validated, "the service must retain its authorization") + require.NoError(t, store.DeleteService(ctx, "owner", svc.ID)) + require.NoError(t, store.DeleteCustomDomain(ctx, "owner", d.ID)) + continue + } + require.NoError(t, deleteErr) + services, err := store.GetAccountServices(ctx, LockingStrengthNone, "owner") + require.NoError(t, err) + assert.Empty(t, services, "a deleted registration must not leave a new service") + } + }) +} diff --git a/management/server/store/sql_store_dns_record.go b/management/server/store/sql_store_dns_record.go new file mode 100644 index 000000000..8a8fc00a1 --- /dev/null +++ b/management/server/store/sql_store_dns_record.go @@ -0,0 +1,109 @@ +package store + +import ( + "context" + "errors" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) CreateDNSRecord(ctx context.Context, record *records.Record) error { + result := s.db.Create(record) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to create dns record to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to create dns record to store") + } + + return nil +} + +func (s *SqlStore) UpdateDNSRecord(ctx context.Context, record *records.Record) error { + result := s.db.Select("*").Save(record) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update dns record to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to update dns record to store") + } + + return nil +} + +func (s *SqlStore) DeleteDNSRecord(ctx context.Context, accountID, zoneID, recordID string) error { + result := s.db.Delete(&records.Record{}, "account_id = ? AND zone_id = ? AND id = ?", accountID, zoneID, recordID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete dns record from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete dns record from store") + } + + if result.RowsAffected == 0 { + return status.NewDNSRecordNotFoundError(recordID) + } + + return nil +} + +func (s *SqlStore) GetDNSRecordByID(ctx context.Context, lockStrength LockingStrength, accountID, zoneID, recordID string) (*records.Record, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var record *records.Record + result := tx.Where("account_id = ? AND zone_id = ? AND id = ?", accountID, zoneID, recordID).Take(&record) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewDNSRecordNotFoundError(recordID) + } + + log.WithContext(ctx).Errorf("failed to get dns record from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get dns record from store") + } + + return record, nil +} + +func (s *SqlStore) GetZoneDNSRecords(ctx context.Context, lockStrength LockingStrength, accountID, zoneID string) ([]*records.Record, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var recordsList []*records.Record + result := tx.Where("account_id = ? AND zone_id = ?", accountID, zoneID).Find(&recordsList) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get zone dns records from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get zone dns records from store") + } + + return recordsList, nil +} + +func (s *SqlStore) GetZoneDNSRecordsByName(ctx context.Context, lockStrength LockingStrength, accountID, zoneID, name string) ([]*records.Record, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var recordsList []*records.Record + result := tx.Where("account_id = ? AND zone_id = ? AND name = ?", accountID, zoneID, name).Find(&recordsList) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get zone dns records by name from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get zone dns records by name from store") + } + + return recordsList, nil +} + +func (s *SqlStore) DeleteZoneDNSRecords(ctx context.Context, accountID, zoneID string) error { + result := s.db.Delete(&records.Record{}, "account_id = ? AND zone_id = ?", accountID, zoneID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete zone dns records from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete zone dns records from store") + } + + return nil +} diff --git a/management/server/store/sql_store_dns_record_test.go b/management/server/store/sql_store_dns_record_test.go new file mode 100644 index 000000000..045dca1e9 --- /dev/null +++ b/management/server/store/sql_store_dns_record_test.go @@ -0,0 +1,260 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_CreateDNSRecord(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + + err = store.CreateDNSRecord(context.Background(), record) + require.NoError(t, err) + + savedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, accountID, zone.ID, record.ID) + require.NoError(t, err) + require.NotNil(t, savedRecord) + assert.Equal(t, record.ID, savedRecord.ID) + assert.Equal(t, record.Name, savedRecord.Name) + assert.Equal(t, record.Type, savedRecord.Type) + assert.Equal(t, record.Content, savedRecord.Content) + assert.Equal(t, record.TTL, savedRecord.TTL) + assert.Equal(t, zone.ID, savedRecord.ZoneID) +} + +func TestSqlStore_GetDNSRecordByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + err = store.CreateDNSRecord(context.Background(), record) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + zoneID string + recordID string + expectError bool + }{ + { + name: "retrieve existing record", + accountID: accountID, + zoneID: zone.ID, + recordID: record.ID, + expectError: false, + }, + { + name: "retrieve non-existing record", + accountID: accountID, + zoneID: zone.ID, + recordID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty record ID", + accountID: accountID, + zoneID: zone.ID, + recordID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + savedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, tt.accountID, tt.zoneID, tt.recordID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, savedRecord) + } else { + require.NoError(t, err) + require.NotNil(t, savedRecord) + assert.Equal(t, tt.recordID, savedRecord.ID) + } + }) + } +} + +func TestSqlStore_GetZoneDNSRecords(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + recordA := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + err = store.CreateDNSRecord(context.Background(), recordA) + require.NoError(t, err) + + recordAAAA := records.NewRecord(accountID, zone.ID, "ipv6.example.com", records.RecordTypeAAAA, "2001:db8::1", 300) + err = store.CreateDNSRecord(context.Background(), recordAAAA) + require.NoError(t, err) + + recordCNAME := records.NewRecord(accountID, zone.ID, "alias.example.com", records.RecordTypeCNAME, "www.example.com", 300) + err = store.CreateDNSRecord(context.Background(), recordCNAME) + require.NoError(t, err) + + allRecords, err := store.GetZoneDNSRecords(context.Background(), LockingStrengthNone, accountID, zone.ID) + require.NoError(t, err) + require.NotNil(t, allRecords) + assert.Equal(t, 3, len(allRecords)) + + recordIDs := make(map[string]bool) + for _, r := range allRecords { + recordIDs[r.ID] = true + } + assert.True(t, recordIDs[recordA.ID]) + assert.True(t, recordIDs[recordAAAA.ID]) + assert.True(t, recordIDs[recordCNAME.ID]) +} + +func TestSqlStore_GetZoneDNSRecordsByName(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + record1 := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + err = store.CreateDNSRecord(context.Background(), record1) + require.NoError(t, err) + + record2 := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeAAAA, "2001:db8::1", 300) + err = store.CreateDNSRecord(context.Background(), record2) + require.NoError(t, err) + + record3 := records.NewRecord(accountID, zone.ID, "mail.example.com", records.RecordTypeA, "192.168.1.2", 600) + err = store.CreateDNSRecord(context.Background(), record3) + require.NoError(t, err) + + recordsByName, err := store.GetZoneDNSRecordsByName(context.Background(), LockingStrengthNone, accountID, zone.ID, "www.example.com") + require.NoError(t, err) + require.NotNil(t, recordsByName) + assert.Equal(t, 2, len(recordsByName)) + + for _, r := range recordsByName { + assert.Equal(t, "www.example.com", r.Name) + } +} + +func TestSqlStore_UpdateDNSRecord(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + err = store.CreateDNSRecord(context.Background(), record) + require.NoError(t, err) + + record.Name = "api.example.com" + record.Content = "192.168.1.100" + record.TTL = 600 + + err = store.UpdateDNSRecord(context.Background(), record) + require.NoError(t, err) + + updatedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, accountID, zone.ID, record.ID) + require.NoError(t, err) + require.NotNil(t, updatedRecord) + assert.Equal(t, "api.example.com", updatedRecord.Name) + assert.Equal(t, "192.168.1.100", updatedRecord.Content) + assert.Equal(t, 600, updatedRecord.TTL) +} + +func TestSqlStore_DeleteDNSRecord(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + err = store.CreateDNSRecord(context.Background(), record) + require.NoError(t, err) + + err = store.DeleteDNSRecord(context.Background(), accountID, zone.ID, record.ID) + require.NoError(t, err) + + deletedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, accountID, zone.ID, record.ID) + require.Error(t, err) + require.Nil(t, deletedRecord) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) +} + +func TestSqlStore_DeleteZoneDNSRecords(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + record1 := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) + err = store.CreateDNSRecord(context.Background(), record1) + require.NoError(t, err) + + record2 := records.NewRecord(accountID, zone.ID, "mail.example.com", records.RecordTypeA, "192.168.1.2", 600) + err = store.CreateDNSRecord(context.Background(), record2) + require.NoError(t, err) + + allRecords, err := store.GetZoneDNSRecords(context.Background(), LockingStrengthNone, accountID, zone.ID) + require.NoError(t, err) + assert.Equal(t, 2, len(allRecords)) + + err = store.DeleteZoneDNSRecords(context.Background(), accountID, zone.ID) + require.NoError(t, err) + + remainingRecords, err := store.GetZoneDNSRecords(context.Background(), LockingStrengthNone, accountID, zone.ID) + require.NoError(t, err) + assert.Equal(t, 0, len(remainingRecords)) +} diff --git a/management/server/store/sql_store_domain_expiration.go b/management/server/store/sql_store_domain_expiration.go index 525c54084..e7b09f4ad 100644 --- a/management/server/store/sql_store_domain_expiration.go +++ b/management/server/store/sql_store_domain_expiration.go @@ -53,7 +53,10 @@ func customDomainServices(db *gorm.DB, d *domain.Domain) *gorm.DB { // Shared domain validation permits underscores, and older rows may contain // other LIKE metacharacters. escaped := strings.NewReplacer("!", "!!", "%", "!%", "_", "!_").Replace(name) - return db.Model(&rpservice.Service{}).Where( + // Registrations are unique by name, so another account can hold a subdomain + // of this one and serve from it. Its services derive their cluster from that + // account's own registration and are not dependents of this one. + return db.Model(&rpservice.Service{}).Where(accountIDCondition, d.AccountID).Where( "LOWER(domain) IN ? OR LOWER(domain) LIKE ? ESCAPE '!' OR LOWER(domain) LIKE ? ESCAPE '!'", []string{name, name + "."}, "%."+escaped, "%."+escaped+".", ) diff --git a/management/server/store/sql_store_group.go b/management/server/store/sql_store_group.go new file mode 100644 index 000000000..323a45734 --- /dev/null +++ b/management/server/store/sql_store_group.go @@ -0,0 +1,356 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// CreateGroups creates the given list of groups to the database. +// groupUpsertColumns is the explicit allowlist of columns that get updated when +// CreateGroups / UpdateGroups hit a PK conflict. public_id is intentionally +// omitted so a caller passing an entity with the zero value (e.g. an HTTP +// handler-built struct) cannot reset the persisted public_id during an upsert. +// Keep this in sync with the Group schema in management/server/types/group.go. +func groupUpsertColumns() clause.Set { + return clause.AssignmentColumns([]string{ + "account_id", + "name", + "issued", + "integration_ref_id", + "integration_ref_integration_type", + "resources", + }) +} + +func (s *SqlStore) CreateGroups(ctx context.Context, accountID string, groups []*types.Group) error { + if len(groups) == 0 { + return nil + } + + return s.transaction(ctx, func(tx *gorm.DB) error { + result := tx. + Clauses( + clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + Where: clause.Where{Exprs: []clause.Expression{clause.Eq{Column: "groups.account_id", Value: accountID}}}, + DoUpdates: groupUpsertColumns(), + }, + ). + Omit(clause.Associations). + Create(&groups) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save groups to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to save groups to store") + } + + return nil + }) +} + +// UpdateGroups updates the given list of groups to the database. +func (s *SqlStore) UpdateGroups(ctx context.Context, accountID string, groups []*types.Group) error { + if len(groups) == 0 { + return nil + } + + return s.transaction(ctx, func(tx *gorm.DB) error { + result := tx. + Clauses( + clause.OnConflict{ + Columns: []clause.Column{{Name: "id"}}, + Where: clause.Where{Exprs: []clause.Expression{clause.Eq{Column: "groups.account_id", Value: accountID}}}, + DoUpdates: groupUpsertColumns(), + }, + ). + Omit(clause.Associations). + Create(&groups) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save groups to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to save groups to store") + } + + return nil + }) +} + +func (s *SqlStore) GetAccountGroups(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.Group, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var groups []*types.Group + result := tx.Preload(clause.Associations).Find(&groups, accountIDCondition, accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "accountID not found: index lookup failed") + } + log.WithContext(ctx).Errorf("failed to get account groups from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get account groups from the store") + } + + for _, g := range groups { + g.LoadGroupPeers() + } + + return groups, nil +} + +func (s *SqlStore) GetResourceGroups(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) ([]*types.Group, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var groups []*types.Group + + likePattern := `%"ID":"` + resourceID + `"%` + + result := tx. + Preload(clause.Associations). + Where("resources LIKE ?", likePattern). + Find(&groups) + + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, result.Error + } + + for _, g := range groups { + g.LoadGroupPeers() + } + + return groups, nil +} + +func (s *SqlStore) getGroups(ctx context.Context, accountID string) ([]*types.Group, error) { + const query = `SELECT id, account_id, public_id, name, issued, resources, integration_ref_id, integration_ref_integration_type FROM groups WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + groups, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*types.Group, error) { + var g types.Group + var resources []byte + var refID sql.NullInt64 + var refType sql.NullString + err := row.Scan(&g.ID, &g.AccountID, &g.PublicID, &g.Name, &g.Issued, &resources, &refID, &refType) + if err == nil { + if refID.Valid { + g.IntegrationReference.ID = int(refID.Int64) + } + if refType.Valid { + g.IntegrationReference.IntegrationType = refType.String + } + if resources != nil { + _ = json.Unmarshal(resources, &g.Resources) + } else { + g.Resources = []types.Resource{} + } + g.GroupPeers = []types.GroupPeer{} + g.Peers = []string{} + } + return &g, err + }) + if err != nil { + return nil, err + } + return groups, nil +} + +// AddResourceToGroup adds a resource to a group. Method always needs to run n a transaction +func (s *SqlStore) AddResourceToGroup(ctx context.Context, accountId string, groupID string, resource *types.Resource) error { + var group types.Group + result := s.db.Where(accountAndIDQueryCondition, accountId, groupID).Take(&group) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return status.NewGroupNotFoundError(groupID) + } + + return status.Errorf(status.Internal, "issue finding group: %s", result.Error) + } + + for _, res := range group.Resources { + if res.ID == resource.ID { + return nil + } + } + + group.Resources = append(group.Resources, *resource) + + if err := s.db.Save(&group).Error; err != nil { + return status.Errorf(status.Internal, "issue updating group: %s", err) + } + + return nil +} + +// RemoveResourceFromGroup removes a resource from a group. Method always needs to run in a transaction +func (s *SqlStore) RemoveResourceFromGroup(ctx context.Context, accountId string, groupID string, resourceID string) error { + var group types.Group + result := s.db.Where(accountAndIDQueryCondition, accountId, groupID).Take(&group) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return status.NewGroupNotFoundError(groupID) + } + + return status.Errorf(status.Internal, "issue finding group: %s", result.Error) + } + + for i, res := range group.Resources { + if res.ID == resourceID { + group.Resources = append(group.Resources[:i], group.Resources[i+1:]...) + break + } + } + + if err := s.db.Save(&group).Error; err != nil { + return status.Errorf(status.Internal, "issue updating group: %s", err) + } + + return nil +} + +// GetGroupByID retrieves a group by ID and account ID. +func (s *SqlStore) GetGroupByID(ctx context.Context, lockStrength LockingStrength, accountID, groupID string) (*types.Group, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var group *types.Group + result := tx.Preload(clause.Associations).Take(&group, accountAndIDQueryCondition, accountID, groupID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewGroupNotFoundError(groupID) + } + log.WithContext(ctx).Errorf("failed to get group from store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get group from store") + } + + group.LoadGroupPeers() + + return group, nil +} + +// GetGroupByName retrieves a group by name and account ID. +func (s *SqlStore) GetGroupByName(ctx context.Context, lockStrength LockingStrength, accountID, groupName string) (*types.Group, error) { + tx := s.db + + var group types.Group + + // TODO: This fix is accepted for now, but if we need to handle this more frequently + // we may need to reconsider changing the types. + query := tx.Preload(clause.Associations) + + result := query. + Model(&types.Group{}). + Joins("LEFT JOIN group_peers ON group_peers.group_id = groups.id"). + Where("groups.account_id = ? AND groups.name = ?", accountID, groupName). + Group("groups.id"). + Order("COUNT(group_peers.peer_id) DESC"). + Limit(1). + First(&group) + if err := result.Error; err != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewGroupNotFoundError(groupName) + } + log.WithContext(ctx).Errorf("failed to get group by name from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get group by name from store") + } + + group.LoadGroupPeers() + + return &group, nil +} + +// GetGroupsByIDs retrieves groups by their IDs and account ID. +func (s *SqlStore) GetGroupsByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, groupIDs []string) (map[string]*types.Group, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var groups []*types.Group + result := tx.Preload(clause.Associations).Find(&groups, accountAndIDsQueryCondition, accountID, groupIDs) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get groups by ID's from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get groups by ID's from store") + } + + groupsMap := make(map[string]*types.Group) + for _, group := range groups { + group.LoadGroupPeers() + groupsMap[group.ID] = group + } + + return groupsMap, nil +} + +// CreateGroup creates a group in the store. +func (s *SqlStore) CreateGroup(ctx context.Context, group *types.Group) error { + if group == nil { + return status.Errorf(status.InvalidArgument, "group is nil") + } + + if err := s.db.Omit(clause.Associations).Create(group).Error; err != nil { + log.WithContext(ctx).Errorf("failed to save group to store: %v", err) + return status.Errorf(status.Internal, "failed to save group to store") + } + + return nil +} + +// UpdateGroup updates a group in the store. +func (s *SqlStore) UpdateGroup(ctx context.Context, group *types.Group) error { + if group == nil { + return status.Errorf(status.InvalidArgument, "group is nil") + } + + if err := s.db.Omit(clause.Associations, "public_id").Save(group).Error; err != nil { + log.WithContext(ctx).Errorf("failed to save group to store: %v", err) + return status.Errorf(status.Internal, "failed to save group to store") + } + + return nil +} + +// DeleteGroup deletes a group from the database. +func (s *SqlStore) DeleteGroup(ctx context.Context, accountID, groupID string) error { + result := s.db.Select(clause.Associations). + Delete(&types.Group{}, accountAndIDQueryCondition, accountID, groupID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to delete group from store: %s", result.Error) + return status.Errorf(status.Internal, "failed to delete group from store") + } + + if result.RowsAffected == 0 { + return status.NewGroupNotFoundError(groupID) + } + + return nil +} + +// DeleteGroups deletes groups from the database. +func (s *SqlStore) DeleteGroups(ctx context.Context, accountID string, groupIDs []string) error { + result := s.db.Select(clause.Associations). + Delete(&types.Group{}, accountAndIDsQueryCondition, accountID, groupIDs) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete groups from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete groups from store") + } + + return nil +} diff --git a/management/server/store/sql_store_group_peer.go b/management/server/store/sql_store_group_peer.go new file mode 100644 index 000000000..cc4c44ad1 --- /dev/null +++ b/management/server/store/sql_store_group_peer.go @@ -0,0 +1,229 @@ +package store + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getGroupPeers(ctx context.Context, groupIDs []string) ([]types.GroupPeer, error) { + if len(groupIDs) == 0 { + return nil, nil + } + const query = `SELECT account_id, group_id, peer_id FROM group_peers WHERE group_id = ANY($1)` + rows, err := s.pgxPool().Query(ctx, query, groupIDs) + if err != nil { + return nil, err + } + groupPeers, err := pgx.CollectRows(rows, pgx.RowToStructByName[types.GroupPeer]) + if err != nil { + return nil, err + } + return groupPeers, nil +} + +// AddPeerToAllGroup adds a peer to the 'All' group. Method always needs to run in a transaction +func (s *SqlStore) AddPeerToAllGroup(ctx context.Context, accountID string, peerID string) error { + var groupID string + _ = s.db.Model(types.Group{}). + Select("id"). + Where("account_id = ? AND name = ?", accountID, "All"). + Limit(1). + Scan(&groupID) + + if groupID == "" { + return status.Errorf(status.NotFound, "group 'All' not found for account %s", accountID) + } + + err := s.db.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "group_id"}, {Name: "peer_id"}}, + DoNothing: true, + }).Create(&types.GroupPeer{ + AccountID: accountID, + GroupID: groupID, + PeerID: peerID, + }).Error + if err != nil { + return status.Errorf(status.Internal, "error adding peer to group 'All': %v", err) + } + + return nil +} + +// AddPeerToGroup adds a peer to a group +func (s *SqlStore) AddPeerToGroup(ctx context.Context, accountID, peerID, groupID string) error { + peer := &types.GroupPeer{ + AccountID: accountID, + GroupID: groupID, + PeerID: peerID, + } + + err := s.db.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "group_id"}, {Name: "peer_id"}}, + DoNothing: true, + }).Create(peer).Error + if err != nil { + log.WithContext(ctx).Errorf("failed to add peer %s to group %s for account %s: %v", peerID, groupID, accountID, err) + return status.Errorf(status.Internal, "failed to add peer to group") + } + + return nil +} + +// RemovePeerFromGroup removes a peer from a group +func (s *SqlStore) RemovePeerFromGroup(ctx context.Context, peerID string, groupID string) error { + err := s.db. + Delete(&types.GroupPeer{}, "group_id = ? AND peer_id = ?", groupID, peerID).Error + if err != nil { + log.WithContext(ctx).Errorf("failed to remove peer %s from group %s: %v", peerID, groupID, err) + return status.Errorf(status.Internal, "failed to remove peer from group") + } + + return nil +} + +// RemovePeerFromAllGroups removes a peer from all groups +func (s *SqlStore) RemovePeerFromAllGroups(ctx context.Context, peerID string) error { + err := s.db. + Delete(&types.GroupPeer{}, "peer_id = ?", peerID).Error + if err != nil { + log.WithContext(ctx).Errorf("failed to remove peer %s from all groups: %v", peerID, err) + return status.Errorf(status.Internal, "failed to remove peer from all groups") + } + + return nil +} + +// GetPeerGroups retrieves all groups assigned to a specific peer in a given account. +func (s *SqlStore) GetPeerGroups(ctx context.Context, lockStrength LockingStrength, accountId string, peerId string) ([]*types.Group, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var groups []*types.Group + query := tx. + Joins("JOIN group_peers ON group_peers.group_id = groups.id"). + Where("groups.account_id = ? AND group_peers.peer_id = ?", accountId, peerId). + Preload(clause.Associations). + Find(&groups) + + if query.Error != nil { + return nil, query.Error + } + + for _, group := range groups { + group.LoadGroupPeers() + } + + return groups, nil +} + +// GetPeerGroupIDs retrieves all group IDs assigned to a specific peer in a given account. +func (s *SqlStore) GetPeerGroupIDs(ctx context.Context, lockStrength LockingStrength, accountId string, peerId string) ([]string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var groupIDs []string + query := tx. + Model(&types.GroupPeer{}). + Where("account_id = ? AND peer_id = ?", accountId, peerId). + Pluck("group_id", &groupIDs) + + if query.Error != nil { + if errors.Is(query.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "no groups found for peer %s in account %s", peerId, accountId) + } + log.WithContext(ctx).Errorf("failed to get group IDs for peer %s in account %s: %v", peerId, accountId, query.Error) + return nil, status.Errorf(status.Internal, "failed to get group IDs for peer from store") + } + + return groupIDs, nil +} + +func (s *SqlStore) GetAccountGroupPeers(ctx context.Context, lockStrength LockingStrength, accountID string) (map[string]map[string]struct{}, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peers []types.GroupPeer + result := tx.Find(&peers, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get account group peers from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get account group peers from store") + } + + groupPeers := make(map[string]map[string]struct{}) + for _, peer := range peers { + if _, exists := groupPeers[peer.GroupID]; !exists { + groupPeers[peer.GroupID] = make(map[string]struct{}) + } + groupPeers[peer.GroupID][peer.PeerID] = struct{}{} + } + + return groupPeers, nil +} + +func (s *SqlStore) GetPeersByGroupIDs(ctx context.Context, accountID string, groupIDs []string) ([]*nbpeer.Peer, error) { + if len(groupIDs) == 0 { + return []*nbpeer.Peer{}, nil + } + + var peers []*nbpeer.Peer + peerIDsSubquery := s.db.Model(&types.GroupPeer{}). + Select("DISTINCT peer_id"). + Where("account_id = ? AND group_id IN ?", accountID, groupIDs) + + result := s.db.Where("account_id = ? AND id IN (?)", accountID, peerIDsSubquery).Find(&peers) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get peers by group IDs: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get peers by group IDs") + } + + return peers, nil +} + +func (s *SqlStore) GetPeerIDsByGroups(ctx context.Context, accountID string, groupIDs []string) ([]string, error) { + if len(groupIDs) == 0 { + return nil, nil + } + + var peerIDs []string + result := s.db.Model(&types.GroupPeer{}). + Select("DISTINCT peer_id"). + Where("account_id = ? AND group_id IN ?", accountID, groupIDs). + Pluck("peer_id", &peerIDs) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "failed to get peer IDs by groups: %s", result.Error) + } + + return peerIDs, nil +} + +func (s *SqlStore) GetGroupIDsByPeerIDs(ctx context.Context, accountID string, peerIDs []string) ([]string, error) { + if len(peerIDs) == 0 { + return nil, nil + } + + var groupIDs []string + result := s.db.Model(&types.GroupPeer{}). + Select("DISTINCT group_id"). + Where("account_id = ? AND peer_id IN ?", accountID, peerIDs). + Pluck("group_id", &groupIDs) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "failed to get group IDs by peers: %s", result.Error) + } + + return groupIDs, nil +} diff --git a/management/server/store/sql_store_group_peer_test.go b/management/server/store/sql_store_group_peer_test.go new file mode 100644 index 000000000..9f9a9e483 --- /dev/null +++ b/management/server/store/sql_store_group_peer_test.go @@ -0,0 +1,210 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" +) + +func TestSqlStore_AddPeerToGroup(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + peerID := "cfefqs706sqkneg59g4g" + groupID := "cfefqs706sqkneg59g4h" + + group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.NoError(t, err, "failed to get group") + require.Len(t, group.Peers, 0, "group should have 0 peers") + + err = store.AddPeerToGroup(context.Background(), accountID, peerID, groupID) + require.NoError(t, err, "failed to add peer to group") + + group, err = store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.NoError(t, err, "failed to get group") + require.Len(t, group.Peers, 1, "group should have 1 peers") + require.Contains(t, group.Peers, peerID) +} + +func TestSqlStore_AddPeerToAllGroup(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + groupID := "cfefqs706sqkneg59g3g" + + peer := &nbpeer.Peer{ + ID: "peer1", + AccountID: accountID, + DNSLabel: "peer1.domain.test", + } + + group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.NoError(t, err, "failed to get group") + require.Len(t, group.Peers, 2, "group should have 2 peers") + require.NotContains(t, group.Peers, peer.ID) + + err = store.AddPeerToAccount(context.Background(), peer) + require.NoError(t, err, "failed to add peer to account") + + err = store.AddPeerToAllGroup(context.Background(), accountID, peer.ID) + require.NoError(t, err, "failed to add peer to all group") + + group, err = store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.NoError(t, err, "failed to get group") + require.Len(t, group.Peers, 3, "group should have peers") + require.Contains(t, group.Peers, peer.ID) +} + +func TestSqlStore_GetPeerGroups(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + peerID := "cfefqs706sqkneg59g4g" + + groups, err := store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peerID) + require.NoError(t, err) + assert.Len(t, groups, 1) + assert.Equal(t, groups[0].Name, "All") + + err = store.AddPeerToGroup(context.Background(), accountID, peerID, "cfefqs706sqkneg59g4h") + require.NoError(t, err) + + groups, err = store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peerID) + require.NoError(t, err) + assert.Len(t, groups, 2) + + foreignPeerID := "foreign-peer" + err = store.AddPeerToGroup(context.Background(), accountID, foreignPeerID, "cfefqs706sqkneg59g4h") + require.NoError(t, err) + + groups, err = store.GetPeerGroups(context.Background(), LockingStrengthNone, "other-account", foreignPeerID) + require.NoError(t, err) + assert.Empty(t, groups, "groups of another account must not be returned") +} + +func TestSqlStore_GetPeersByGroupIDs(t *testing.T) { + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + group1ID := "test-group-1" + group2ID := "test-group-2" + emptyGroupID := "empty-group" + + peer1 := "cfefqs706sqkneg59g4g" + peer2 := "cfeg6sf06sqkneg59g50" + + tests := []struct { + name string + groupIDs []string + expectedPeers []string + expectedCount int + }{ + { + name: "retrieve peers from single group with multiple peers", + groupIDs: []string{group1ID}, + expectedPeers: []string{peer1, peer2}, + expectedCount: 2, + }, + { + name: "retrieve peers from single group with one peer", + groupIDs: []string{group2ID}, + expectedPeers: []string{peer1}, + expectedCount: 1, + }, + { + name: "retrieve peers from multiple groups (with overlap)", + groupIDs: []string{group1ID, group2ID}, + expectedPeers: []string{peer1, peer2}, // should deduplicate + expectedCount: 2, + }, + { + name: "retrieve peers from existing 'All' group", + groupIDs: []string{"cfefqs706sqkneg59g3g"}, // All group from test data + expectedPeers: []string{peer1, peer2}, + expectedCount: 2, + }, + { + name: "retrieve peers from empty group", + groupIDs: []string{emptyGroupID}, + expectedPeers: []string{}, + expectedCount: 0, + }, + { + name: "retrieve peers from non-existing group", + groupIDs: []string{"non-existing-group"}, + expectedPeers: []string{}, + expectedCount: 0, + }, + { + name: "empty group IDs list", + groupIDs: []string{}, + expectedPeers: []string{}, + expectedCount: 0, + }, + { + name: "mix of existing and non-existing groups", + groupIDs: []string{group1ID, "non-existing-group"}, + expectedPeers: []string{peer1, peer2}, + expectedCount: 2, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + ctx := context.Background() + + groups := []*types.Group{ + { + ID: group1ID, + AccountID: accountID, + }, + { + ID: group2ID, + AccountID: accountID, + }, + } + require.NoError(t, store.CreateGroups(ctx, accountID, groups)) + + otherAccount := newAccountWithId(ctx, "other-account", "other-user", "") + require.NoError(t, store.SaveAccount(ctx, otherAccount)) + foreignPeer := &nbpeer.Peer{ID: "foreign-peer", AccountID: otherAccount.Id} + require.NoError(t, store.AddPeerToAccount(ctx, foreignPeer)) + + require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer1, group1ID)) + require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer2, group1ID)) + require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer1, group2ID)) + require.NoError(t, store.AddPeerToGroup(ctx, accountID, foreignPeer.ID, group1ID)) + + peers, err := store.GetPeersByGroupIDs(ctx, accountID, tt.groupIDs) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + + if tt.expectedCount > 0 { + actualPeerIDs := make([]string, len(peers)) + for i, peer := range peers { + actualPeerIDs[i] = peer.ID + } + assert.ElementsMatch(t, tt.expectedPeers, actualPeerIDs) + + // Verify all returned peers belong to the correct account + for _, peer := range peers { + assert.Equal(t, accountID, peer.AccountID) + } + } + }) + } +} diff --git a/management/server/store/sql_store_group_test.go b/management/server/store/sql_store_group_test.go new file mode 100644 index 000000000..58291ab44 --- /dev/null +++ b/management/server/store/sql_store_group_test.go @@ -0,0 +1,289 @@ +package store + +import ( + "context" + "fmt" + "os" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlite_GetGroupByName(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + if err != nil { + t.Fatal(err) + } + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + group, err := store.GetGroupByName(context.Background(), LockingStrengthNone, accountID, "All") + require.NoError(t, err) + require.True(t, group.IsGroupAll()) +} + +func TestSqlStore_GetGroupsByIDs(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + groupIDs []string + expectedCount int + }{ + { + name: "retrieve existing groups by existing IDs", + groupIDs: []string{"cfefqs706sqkneg59g4g", "cfefqs706sqkneg59g3g"}, + expectedCount: 2, + }, + { + name: "empty group IDs list", + groupIDs: []string{}, + expectedCount: 0, + }, + { + name: "non-existing group IDs", + groupIDs: []string{"nonexistent1", "nonexistent2"}, + expectedCount: 0, + }, + { + name: "mixed existing and non-existing group IDs", + groupIDs: []string{"cfefqs706sqkneg59g4g", "nonexistent"}, + expectedCount: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + groups, err := store.GetGroupsByIDs(context.Background(), LockingStrengthNone, accountID, tt.groupIDs) + require.NoError(t, err) + require.Len(t, groups, tt.expectedCount) + }) + } +} + +func TestSqlStore_CreateGroup(t *testing.T) { + if os.Getenv("CI") == "true" { + t.Log("Skipping MySQL test on CI") + } + t.Setenv("NETBIRD_STORE_ENGINE", string(types.MysqlStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + group := &types.Group{ + ID: "group-id", + AccountID: accountID, + Issued: "api", + Peers: []string{}, + Resources: []types.Resource{}, + GroupPeers: []types.GroupPeer{}, + } + err = store.CreateGroup(context.Background(), group) + require.NoError(t, err) + + savedGroup, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, "group-id") + require.NoError(t, err) + require.Equal(t, savedGroup, group) +} + +func TestSqlStore_CreateUpdateGroups(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + groups := []*types.Group{ + { + ID: "group-1", + AccountID: accountID, + Issued: "api", + Peers: []string{}, + Resources: []types.Resource{}, + GroupPeers: []types.GroupPeer{}, + }, + { + ID: "group-2", + AccountID: accountID, + Issued: "integration", + Peers: []string{}, + Resources: []types.Resource{}, + GroupPeers: []types.GroupPeer{}, + }, + } + err = store.CreateGroups(context.Background(), accountID, groups) + require.NoError(t, err) + + groups[1].Peers = []string{} + err = store.UpdateGroups(context.Background(), accountID, groups) + require.NoError(t, err) + + group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groups[1].ID) + require.NoError(t, err) + require.Equal(t, groups[1], group) +} + +func TestSqlStore_DeleteGroup(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + groupID string + expectError bool + }{ + { + name: "delete existing group", + groupID: "cfefqs706sqkneg59g4g", + expectError: false, + }, + { + name: "delete non-existing group", + groupID: "non-existing-group-id", + expectError: true, + }, + { + name: "delete with empty group ID", + groupID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := store.DeleteGroup(context.Background(), accountID, tt.groupID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + } else { + require.NoError(t, err) + + group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, tt.groupID) + require.Error(t, err) + require.Nil(t, group) + } + }) + } +} + +func TestSqlStore_DeleteGroups(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + groupIDs []string + expectError bool + }{ + { + name: "delete multiple existing groups", + groupIDs: []string{"cfefqs706sqkneg59g4g", "cfefqs706sqkneg59g3g"}, + expectError: false, + }, + { + name: "delete non-existing groups", + groupIDs: []string{"non-existing-id-1", "non-existing-id-2"}, + expectError: false, + }, + { + name: "delete with empty group IDs list", + groupIDs: []string{}, + expectError: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := store.DeleteGroups(context.Background(), accountID, tt.groupIDs) + if tt.expectError { + require.Error(t, err) + } else { + require.NoError(t, err) + + for _, groupID := range tt.groupIDs { + group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.Error(t, err) + require.Nil(t, group) + } + } + }) + } +} + +func TestSqlStore_AddAndRemoveResourceFromGroup(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + require.NoError(t, err) + t.Cleanup(cleanup) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + resourceId := "ctc4nci7qv9061u6ilfg" + groupID := "cs1tnh0hhcjnqoiuebeg" + + res := &types.Resource{ + ID: resourceId, + Type: "host", + } + err = store.AddResourceToGroup(context.Background(), accountID, groupID, res) + require.NoError(t, err) + + group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.NoError(t, err) + require.Contains(t, group.Resources, *res) + + groups, err := store.GetResourceGroups(context.Background(), LockingStrengthNone, accountID, resourceId) + require.NoError(t, err) + require.Len(t, groups, 1) + + err = store.RemoveResourceFromGroup(context.Background(), accountID, groupID, res.ID) + require.NoError(t, err) + + group, err = store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) + require.NoError(t, err) + require.NotContains(t, group.Resources, *res) +} + +func TestSqlStore_SaveGroups_LargeBatch(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + accountGroups, err := store.GetAccountGroups(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Len(t, accountGroups, 3) + + groupsToSave := make([]*types.Group, 0) + + for i := 1; i <= 8000; i++ { + groupsToSave = append(groupsToSave, &types.Group{ + ID: fmt.Sprintf("%d", i), + AccountID: accountID, + Name: fmt.Sprintf("group-%d", i), + }) + } + + err = store.CreateGroups(context.Background(), accountID, groupsToSave) + require.NoError(t, err) + + accountGroups, err = store.GetAccountGroups(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Equal(t, 8003, len(accountGroups)) +} diff --git a/management/server/store/sql_store_idp_migration.go b/management/server/store/sql_store_idp_migration.go index 64962845b..760d2967c 100644 --- a/management/server/store/sql_store_idp_migration.go +++ b/management/server/store/sql_store_idp_migration.go @@ -57,11 +57,11 @@ func (s *SqlStore) ListUsers(ctx context.Context) ([]*types.User, error) { // txDeferFKConstraints defers foreign key constraint checks for the duration of the transaction. // MySQL is already handled by s.transaction (SET FOREIGN_KEY_CHECKS = 0). func (s *SqlStore) txDeferFKConstraints(tx *gorm.DB) error { - if s.storeEngine == types.SqliteStoreEngine { + if s.conn.Engine() == types.SqliteStoreEngine { return tx.Exec("PRAGMA defer_foreign_keys = ON").Error } - if s.storeEngine != types.PostgresStoreEngine { + if s.conn.Engine() != types.PostgresStoreEngine { return nil } @@ -86,7 +86,7 @@ func (s *SqlStore) txDeferFKConstraints(tx *gorm.DB) error { // txRestoreFKConstraints reverts FK constraints back to NOT DEFERRABLE after the // deferred updates are done but before the transaction commits. func (s *SqlStore) txRestoreFKConstraints(tx *gorm.DB) error { - if s.storeEngine != types.PostgresStoreEngine { + if s.conn.Engine() != types.PostgresStoreEngine { return nil } @@ -138,7 +138,7 @@ func (s *SqlStore) UpdateUserID(ctx context.Context, accountID, oldUserID, newUs } log.Info("Updating user ID in the store") - err := s.transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { if err := s.txDeferFKConstraints(tx); err != nil { return err } @@ -161,7 +161,7 @@ func (s *SqlStore) UpdateUserID(ctx context.Context, accountID, oldUserID, newUs } log.Info("Restoring FK constraints") - err = s.transaction(func(tx *gorm.DB) error { + err = s.transaction(ctx, func(tx *gorm.DB) error { if err := s.txRestoreFKConstraints(tx); err != nil { return fmt.Errorf("restore FK constraints: %w", err) } diff --git a/management/server/store/sql_store_installation.go b/management/server/store/sql_store_installation.go new file mode 100644 index 000000000..2bdfb9af1 --- /dev/null +++ b/management/server/store/sql_store_installation.go @@ -0,0 +1,29 @@ +package store + +import ( + "context" + + "gorm.io/gorm/clause" +) + +type installation struct { + ID uint `gorm:"primaryKey"` + InstallationIDValue string +} + +func (s *SqlStore) SaveInstallationID(_ context.Context, ID string) error { + installation := installation{InstallationIDValue: ID} + installation.ID = uint(s.installationPK) + + return s.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&installation).Error +} + +func (s *SqlStore) GetInstallationID() string { + var installation installation + + if result := s.db.Take(&installation, idQueryCondition, s.installationPK); result.Error != nil { + return "" + } + + return installation.InstallationIDValue +} diff --git a/management/server/store/sql_store_job.go b/management/server/store/sql_store_job.go new file mode 100644 index 000000000..b5cc6c603 --- /dev/null +++ b/management/server/store/sql_store_job.go @@ -0,0 +1,103 @@ +package store + +import ( + "context" + "errors" + "time" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// SaveJob persists a job in DB +func (s *SqlStore) CreatePeerJob(ctx context.Context, job *types.Job) error { + result := s.db.Create(job) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to create job in store: %s", result.Error) + return status.Errorf(status.Internal, "failed to create job in store") + } + return nil +} + +func (s *SqlStore) CompletePeerJob(ctx context.Context, job *types.Job) error { + result := s.db. + Model(&types.Job{}). + Where(idQueryCondition, job.ID). + Updates(job) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update job in store: %s", result.Error) + return status.Errorf(status.Internal, "failed to update job in store") + } + return nil +} + +// job was pending for too long and has been cancelled +func (s *SqlStore) MarkPendingJobsAsFailed(ctx context.Context, accountID, peerID, jobID, reason string) error { + now := time.Now().UTC() + result := s.db. + Model(&types.Job{}). + Where(accountAndPeerIDQueryCondition+" AND id = ?"+" AND status = ?", accountID, peerID, jobID, types.JobStatusPending). + Updates(types.Job{ + Status: types.JobStatusFailed, + FailedReason: reason, + CompletedAt: &now, + }) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to mark pending jobs as Failed job in store: %s", result.Error) + return status.Errorf(status.Internal, "failed to mark pending job as Failed in store") + } + return nil +} + +// job was pending for too long and has been cancelled +func (s *SqlStore) MarkAllPendingJobsAsFailed(ctx context.Context, accountID, peerID, reason string) error { + now := time.Now().UTC() + result := s.db. + Model(&types.Job{}). + Where(accountAndPeerIDQueryCondition+" AND status = ?", accountID, peerID, types.JobStatusPending). + Updates(types.Job{ + Status: types.JobStatusFailed, + FailedReason: reason, + CompletedAt: &now, + }) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to mark pending jobs as Failed job in store: %s", result.Error) + return status.Errorf(status.Internal, "failed to mark pending job as Failed in store") + } + return nil +} + +// GetJobByID fetches job by ID +func (s *SqlStore) GetPeerJobByID(ctx context.Context, accountID, jobID string) (*types.Job, error) { + var job types.Job + err := s.db. + Where(accountAndIDQueryCondition, accountID, jobID). + First(&job).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "job %s not found", jobID) + } + if err != nil { + log.WithContext(ctx).Errorf("failed to fetch job from store: %s", err) + return nil, err + } + return &job, nil +} + +// get all jobs +func (s *SqlStore) GetPeerJobs(ctx context.Context, accountID, peerID string) ([]*types.Job, error) { + var jobs []*types.Job + err := s.db. + Where(accountAndPeerIDQueryCondition, accountID, peerID). + Order("created_at DESC"). + Find(&jobs).Error + if err != nil { + log.WithContext(ctx).Errorf("failed to fetch jobs from store: %s", err) + return nil, err + } + + return jobs, nil +} diff --git a/management/server/store/sql_store_name_server_group.go b/management/server/store/sql_store_name_server_group.go new file mode 100644 index 000000000..595921913 --- /dev/null +++ b/management/server/store/sql_store_name_server_group.go @@ -0,0 +1,124 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getNameServerGroups(ctx context.Context, accountID string) ([]nbdns.NameServerGroup, error) { + const query = `SELECT id, account_id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled FROM name_server_groups WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + nsgs, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nbdns.NameServerGroup, error) { + var n nbdns.NameServerGroup + var ns, groups, domains []byte + var primary, enabled, searchDomainsEnabled sql.NullBool + err := row.Scan(&n.ID, &n.AccountID, &n.PublicID, &n.Name, &n.Description, &ns, &groups, &primary, &domains, &enabled, &searchDomainsEnabled) + if err == nil { + if primary.Valid { + n.Primary = primary.Bool + } + if enabled.Valid { + n.Enabled = enabled.Bool + } + if searchDomainsEnabled.Valid { + n.SearchDomainsEnabled = searchDomainsEnabled.Bool + } + if ns != nil { + _ = json.Unmarshal(ns, &n.NameServers) + } else { + n.NameServers = []nbdns.NameServer{} + } + if groups != nil { + _ = json.Unmarshal(groups, &n.Groups) + } else { + n.Groups = []string{} + } + if domains != nil { + _ = json.Unmarshal(domains, &n.Domains) + } else { + n.Domains = []string{} + } + } + return n, err + }) + if err != nil { + return nil, err + } + return nsgs, nil +} + +// GetAccountNameServerGroups retrieves name server groups for an account. +func (s *SqlStore) GetAccountNameServerGroups(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbdns.NameServerGroup, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var nsGroups []*nbdns.NameServerGroup + result := tx.Find(&nsGroups, accountIDCondition, accountID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get name server groups from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get name server groups from store") + } + + return nsGroups, nil +} + +// GetNameServerGroupByID retrieves a name server group by its ID and account ID. +func (s *SqlStore) GetNameServerGroupByID(ctx context.Context, lockStrength LockingStrength, accountID, nsGroupID string) (*nbdns.NameServerGroup, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var nsGroup *nbdns.NameServerGroup + result := tx. + Take(&nsGroup, accountAndIDQueryCondition, accountID, nsGroupID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewNameServerGroupNotFoundError(nsGroupID) + } + log.WithContext(ctx).Errorf("failed to get name server group from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get name server group from store") + } + + return nsGroup, nil +} + +// SaveNameServerGroup saves a name server group to the database. +func (s *SqlStore) SaveNameServerGroup(ctx context.Context, nameServerGroup *nbdns.NameServerGroup) error { + result := s.db.Save(nameServerGroup) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to save name server group to the store: %s", err) + return status.Errorf(status.Internal, "failed to save name server group to store") + } + return nil +} + +// DeleteNameServerGroup deletes a name server group from the database. +func (s *SqlStore) DeleteNameServerGroup(ctx context.Context, accountID, nsGroupID string) error { + result := s.db.Delete(&nbdns.NameServerGroup{}, accountAndIDQueryCondition, accountID, nsGroupID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to delete name server group from the store: %s", err) + return status.Errorf(status.Internal, "failed to delete name server group from store") + } + + if result.RowsAffected == 0 { + return status.NewNameServerGroupNotFoundError(nsGroupID) + } + + return nil +} diff --git a/management/server/store/sql_store_name_server_group_test.go b/management/server/store/sql_store_name_server_group_test.go new file mode 100644 index 000000000..ae849c383 --- /dev/null +++ b/management/server/store/sql_store_name_server_group_test.go @@ -0,0 +1,143 @@ +package store + +import ( + "context" + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + + nbdns "github.com/netbirdio/netbird/dns" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetAccountNameServerGroups(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectedCount int + }{ + { + name: "retrieve name server groups by existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectedCount: 1, + }, + { + name: "non-existing account ID", + accountID: "nonexistent", + expectedCount: 0, + }, + { + name: "empty account ID", + accountID: "", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetAccountNameServerGroups(context.Background(), LockingStrengthNone, tt.accountID) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } + +} + +func TestSqlStore_GetNameServerByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + nsGroupID string + expectError bool + }{ + { + name: "retrieve existing nameserver group", + nsGroupID: "csqdelq7qv97ncu7d9t0", + expectError: false, + }, + { + name: "retrieve non-existing nameserver group", + nsGroupID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty nameserver group ID", + nsGroupID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + nsGroup, err := store.GetNameServerGroupByID(context.Background(), LockingStrengthNone, accountID, tt.nsGroupID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, nsGroup) + } else { + require.NoError(t, err) + require.NotNil(t, nsGroup) + require.Equal(t, tt.nsGroupID, nsGroup.ID) + } + }) + } +} + +func TestSqlStore_SaveNameServerGroup(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + nsGroup := &nbdns.NameServerGroup{ + ID: "ns-group-id", + AccountID: accountID, + Name: "NS Group", + NameServers: []nbdns.NameServer{ + { + IP: netip.MustParseAddr("8.8.8.8"), + NSType: 1, + Port: 53, + }, + }, + Groups: []string{"groupA"}, + Primary: true, + Enabled: true, + SearchDomainsEnabled: false, + } + + err = store.SaveNameServerGroup(context.Background(), nsGroup) + require.NoError(t, err) + + saveNSGroup, err := store.GetNameServerGroupByID(context.Background(), LockingStrengthNone, accountID, nsGroup.ID) + require.NoError(t, err) + require.Equal(t, saveNSGroup, nsGroup) +} + +func TestSqlStore_DeleteNameServerGroup(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + nsGroupID := "csqdelq7qv97ncu7d9t0" + + err = store.DeleteNameServerGroup(context.Background(), accountID, nsGroupID) + require.NoError(t, err) + + nsGroup, err := store.GetNameServerGroupByID(context.Background(), LockingStrengthNone, accountID, nsGroupID) + require.Error(t, err) + require.Nil(t, nsGroup) +} diff --git a/management/server/store/sql_store_network.go b/management/server/store/sql_store_network.go new file mode 100644 index 000000000..66333a4fb --- /dev/null +++ b/management/server/store/sql_store_network.go @@ -0,0 +1,91 @@ +package store + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + networkTypes "github.com/netbirdio/netbird/management/server/networks/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) { + const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + networks, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkTypes.Network]) + if err != nil { + return nil, err + } + result := make([]*networkTypes.Network, len(networks)) + for i := range networks { + result[i] = &networks[i] + } + return result, nil +} + +func (s *SqlStore) GetAccountNetworks(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*networkTypes.Network, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var networks []*networkTypes.Network + result := tx.Find(&networks, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get networks from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get networks from store") + } + + return networks, nil +} + +func (s *SqlStore) GetNetworkByID(ctx context.Context, lockStrength LockingStrength, accountID, networkID string) (*networkTypes.Network, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var network *networkTypes.Network + result := tx.Take(&network, accountAndIDQueryCondition, accountID, networkID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewNetworkNotFoundError(networkID) + } + + log.WithContext(ctx).Errorf("failed to get network from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network from store") + } + + return network, nil +} + +func (s *SqlStore) SaveNetwork(ctx context.Context, network *networkTypes.Network) error { + result := s.db.Save(network) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save network to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to save network to store") + } + + return nil +} + +func (s *SqlStore) DeleteNetwork(ctx context.Context, accountID, networkID string) error { + result := s.db.Delete(&networkTypes.Network{}, accountAndIDQueryCondition, accountID, networkID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete network from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete network from store") + } + + if result.RowsAffected == 0 { + return status.NewNetworkNotFoundError(networkID) + } + + return nil +} diff --git a/management/server/store/sql_store_network_resource.go b/management/server/store/sql_store_network_resource.go new file mode 100644 index 000000000..dd4352ea0 --- /dev/null +++ b/management/server/store/sql_store_network_resource.go @@ -0,0 +1,167 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getNetworkResources(ctx context.Context, accountID string) ([]*resourceTypes.NetworkResource, error) { + const query = `SELECT id, network_id, account_id, public_id, name, description, type, domain, prefix, enabled FROM network_resources WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + resources, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (resourceTypes.NetworkResource, error) { + var r resourceTypes.NetworkResource + var prefix []byte + var enabled sql.NullBool + err := row.Scan(&r.ID, &r.NetworkID, &r.AccountID, &r.PublicID, &r.Name, &r.Description, &r.Type, &r.Domain, &prefix, &enabled) + if err == nil { + if enabled.Valid { + r.Enabled = enabled.Bool + } + if prefix != nil { + _ = json.Unmarshal(prefix, &r.Prefix) + } + } + return r, err + }) + if err != nil { + return nil, err + } + result := make([]*resourceTypes.NetworkResource, len(resources)) + for i := range resources { + result[i] = &resources[i] + } + return result, nil +} + +func (s *SqlStore) GetNetworkResourcesByNetID(ctx context.Context, lockStrength LockingStrength, accountID, networkID string) ([]*resourceTypes.NetworkResource, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netResources []*resourceTypes.NetworkResource + result := tx. + Find(&netResources, "account_id = ? AND network_id = ?", accountID, networkID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get network resources from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network resources from store") + } + + return netResources, nil +} + +func (s *SqlStore) GetNetworkResourcesByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*resourceTypes.NetworkResource, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netResources []*resourceTypes.NetworkResource + result := tx. + Find(&netResources, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get network resources from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network resources from store") + } + + return netResources, nil +} + +func (s *SqlStore) GetNetworkResourceByID(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) (*resourceTypes.NetworkResource, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netResources *resourceTypes.NetworkResource + result := tx. + Take(&netResources, accountAndIDQueryCondition, accountID, resourceID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewNetworkResourceNotFoundError(resourceID) + } + log.WithContext(ctx).Errorf("failed to get network resource from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network resource from store") + } + + return netResources, nil +} + +// GetNetworkResourceByIDOrPublicID retrieves a network resource by either its ID or its +// PublicID. See GetPolicyByIDOrPublicID for why peer-reported references need both. +func (s *SqlStore) GetNetworkResourceByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) (*resourceTypes.NetworkResource, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netResources *resourceTypes.NetworkResource + result := tx. + Take(&netResources, accountAndAnyIDQueryCondition, accountID, resourceID, resourceID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewNetworkResourceNotFoundError(resourceID) + } + log.WithContext(ctx).Errorf("failed to get network resource from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network resource from store") + } + + return netResources, nil +} + +func (s *SqlStore) GetNetworkResourceByName(ctx context.Context, lockStrength LockingStrength, accountID, resourceName string) (*resourceTypes.NetworkResource, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netResources *resourceTypes.NetworkResource + result := tx. + Take(&netResources, "account_id = ? AND name = ?", accountID, resourceName) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewNetworkResourceNotFoundError(resourceName) + } + log.WithContext(ctx).Errorf("failed to get network resource from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network resource from store") + } + + return netResources, nil +} + +func (s *SqlStore) SaveNetworkResource(ctx context.Context, resource *resourceTypes.NetworkResource) error { + result := s.db.Save(resource) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save network resource to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to save network resource to store") + } + + return nil +} + +func (s *SqlStore) DeleteNetworkResource(ctx context.Context, accountID, resourceID string) error { + result := s.db.Delete(&resourceTypes.NetworkResource{}, accountAndIDQueryCondition, accountID, resourceID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete network resource from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete network resource from store") + } + + if result.RowsAffected == 0 { + return status.NewNetworkResourceNotFoundError(resourceID) + } + + return nil +} diff --git a/management/server/store/sql_store_network_resource_test.go b/management/server/store/sql_store_network_resource_test.go new file mode 100644 index 000000000..620874e54 --- /dev/null +++ b/management/server/store/sql_store_network_resource_test.go @@ -0,0 +1,161 @@ +package store + +import ( + "context" + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + + resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetNetworkResourcesByNetID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + networkID string + expectedCount int + }{ + { + name: "retrieve resources by existing network ID", + networkID: "ct286bi7qv930dsrrug0", + expectedCount: 1, + }, + { + name: "retrieve resources by non-existing network ID", + networkID: "non-existent", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + netResources, err := store.GetNetworkResourcesByNetID(context.Background(), LockingStrengthNone, accountID, tt.networkID) + require.NoError(t, err) + require.Len(t, netResources, tt.expectedCount) + }) + } +} + +func TestSqlStore_GetNetworkResourceByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + netResourceID string + expectError bool + }{ + { + name: "retrieve existing network resource ID", + netResourceID: "ctc4nci7qv9061u6ilfg", + expectError: false, + }, + { + name: "retrieve non-existing network resource ID", + netResourceID: "non-existing", + expectError: true, + }, + { + name: "retrieve network with empty resource ID", + netResourceID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + netResource, err := store.GetNetworkResourceByID(context.Background(), LockingStrengthNone, accountID, tt.netResourceID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, netResource) + } else { + require.NoError(t, err) + require.NotNil(t, netResource) + require.Equal(t, tt.netResourceID, netResource.ID) + } + }) + } +} + +func TestSqlStore_GetNetworkResourceByIDOrPublicID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + netResourceID := "ctc4nci7qv9061u6ilfg" + + netResource, err := store.GetNetworkResourceByID(context.Background(), LockingStrengthNone, accountID, netResourceID) + require.NoError(t, err) + require.NotEmpty(t, netResource.PublicID) + + for _, id := range []string{netResourceID, netResource.PublicID} { + netResource, err := store.GetNetworkResourceByIDOrPublicID(context.Background(), LockingStrengthNone, accountID, id) + require.NoError(t, err) + require.Equal(t, netResourceID, netResource.ID) + } + + netResource, err = store.GetNetworkResourceByIDOrPublicID(context.Background(), LockingStrengthNone, accountID, "non-existing") + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, netResource) +} + +func TestSqlStore_SaveNetworkResource(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + networkID := "ct286bi7qv930dsrrug0" + + netResource, err := resourceTypes.NewNetworkResource(accountID, networkID, "resource-name", "", "example.com", []string{}, true) + require.NoError(t, err) + + err = store.SaveNetworkResource(context.Background(), netResource) + require.NoError(t, err) + + savedNetResource, err := store.GetNetworkResourceByID(context.Background(), LockingStrengthNone, accountID, netResource.ID) + require.NoError(t, err) + require.Equal(t, netResource.ID, savedNetResource.ID) + require.Equal(t, netResource.Name, savedNetResource.Name) + require.Equal(t, netResource.NetworkID, savedNetResource.NetworkID) + require.Equal(t, netResource.Type, resourceTypes.NetworkResourceType("domain")) + require.Equal(t, netResource.Domain, "example.com") + require.Equal(t, netResource.AccountID, savedNetResource.AccountID) + require.Equal(t, netResource.Prefix, netip.Prefix{}) +} + +func TestSqlStore_DeleteNetworkResource(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + netResourceID := "ctc4nci7qv9061u6ilfg" + + err = store.DeleteNetworkResource(context.Background(), accountID, netResourceID) + require.NoError(t, err) + + netResource, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, netResourceID) + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, sErr.Type()) + require.Nil(t, netResource) +} diff --git a/management/server/store/sql_store_network_router.go b/management/server/store/sql_store_network_router.go new file mode 100644 index 000000000..bb5cb6621 --- /dev/null +++ b/management/server/store/sql_store_network_router.go @@ -0,0 +1,208 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + networkTypes "github.com/netbirdio/netbird/management/server/networks/types" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getNetworkRouters(ctx context.Context, accountID string) ([]*routerTypes.NetworkRouter, error) { + const query = `SELECT id, network_id, account_id, public_id, peer, peer_groups, masquerade, metric, enabled FROM network_routers WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + routers, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (routerTypes.NetworkRouter, error) { + var r routerTypes.NetworkRouter + var peerGroups []byte + var masquerade, enabled sql.NullBool + var metric sql.NullInt64 + err := row.Scan(&r.ID, &r.NetworkID, &r.AccountID, &r.PublicID, &r.Peer, &peerGroups, &masquerade, &metric, &enabled) + if err == nil { + if masquerade.Valid { + r.Masquerade = masquerade.Bool + } + if enabled.Valid { + r.Enabled = enabled.Bool + } + if metric.Valid { + r.Metric = int(metric.Int64) + } + if peerGroups != nil { + _ = json.Unmarshal(peerGroups, &r.PeerGroups) + } + } + return r, err + }) + if err != nil { + return nil, err + } + result := make([]*routerTypes.NetworkRouter, len(routers)) + for i := range routers { + result[i] = &routers[i] + } + return result, nil +} + +func (s *SqlStore) GetNetworkRoutersByNetID(ctx context.Context, lockStrength LockingStrength, accountID, netID string) ([]*routerTypes.NetworkRouter, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netRouters []*routerTypes.NetworkRouter + result := tx. + Find(&netRouters, "account_id = ? AND network_id = ?", accountID, netID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get network routers from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network routers from store") + } + + return netRouters, nil +} + +func (s *SqlStore) GetNetworkRoutersByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*routerTypes.NetworkRouter, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netRouters []*routerTypes.NetworkRouter + result := tx. + Find(&netRouters, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get network routers from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network routers from store") + } + + return netRouters, nil +} + +func (s *SqlStore) GetNetworkRouterByID(ctx context.Context, lockStrength LockingStrength, accountID, routerID string) (*routerTypes.NetworkRouter, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var netRouter *routerTypes.NetworkRouter + result := tx. + Take(&netRouter, accountAndIDQueryCondition, accountID, routerID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewNetworkRouterNotFoundError(routerID) + } + log.WithContext(ctx).Errorf("failed to get network router from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get network router from store") + } + + return netRouter, nil +} + +func (s *SqlStore) CreateNetworkRouter(ctx context.Context, router *routerTypes.NetworkRouter) error { + if err := s.db.Create(router).Error; err != nil { + log.WithContext(ctx).Errorf("failed to create network router in store: %v", err) + return status.Errorf(status.Internal, "failed to create network router in store") + } + + return nil +} + +func (s *SqlStore) UpdateNetworkRouter(ctx context.Context, router *routerTypes.NetworkRouter) error { + result := s.db. + Select("*"). + Where(accountAndIDQueryCondition, router.AccountID, router.ID). + Updates(router) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update network router in store: %v", result.Error) + return status.Errorf(status.Internal, "failed to update network router in store") + } + + if result.RowsAffected == 0 { + return status.NewNetworkRouterNotFoundError(router.ID) + } + + return nil +} + +func (s *SqlStore) DeleteNetworkRouter(ctx context.Context, accountID, routerID string) error { + result := s.db.Delete(&routerTypes.NetworkRouter{}, accountAndIDQueryCondition, accountID, routerID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete network router from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete network router from store") + } + + if result.RowsAffected == 0 { + return status.NewNetworkRouterNotFoundError(routerID) + } + + return nil +} + +// GetRoutingPeerNetworks returns the distinct network names where the peer is assigned as a routing peer +// in an enabled network router, either directly or via peer groups. +func (s *SqlStore) GetRoutingPeerNetworks(_ context.Context, accountID, peerID string) ([]string, error) { + var routers []*routerTypes.NetworkRouter + if err := s.db.Select("peer, peer_groups, network_id").Where("account_id = ? AND enabled = true", accountID).Find(&routers).Error; err != nil { + return nil, status.Errorf(status.Internal, "failed to get enabled routers: %v", err) + } + + if len(routers) == 0 { + return nil, nil + } + + var groupPeers []types.GroupPeer + if err := s.db.Select("group_id").Where("account_id = ? AND peer_id = ?", accountID, peerID).Find(&groupPeers).Error; err != nil { + return nil, status.Errorf(status.Internal, "failed to get peer group memberships: %v", err) + } + + groupSet := make(map[string]struct{}, len(groupPeers)) + for _, gp := range groupPeers { + groupSet[gp.GroupID] = struct{}{} + } + + networkIDs := make(map[string]struct{}) + for _, r := range routers { + if r.Peer == peerID { + networkIDs[r.NetworkID] = struct{}{} + } else if r.Peer == "" { + for _, pg := range r.PeerGroups { + if _, ok := groupSet[pg]; ok { + networkIDs[r.NetworkID] = struct{}{} + break + } + } + } + } + + if len(networkIDs) == 0 { + return nil, nil + } + + ids := make([]string, 0, len(networkIDs)) + for id := range networkIDs { + ids = append(ids, id) + } + + var networks []*networkTypes.Network + if err := s.db.Select("name").Where("account_id = ? AND id IN ?", accountID, ids).Find(&networks).Error; err != nil { + return nil, status.Errorf(status.Internal, "failed to get networks: %v", err) + } + + names := make([]string, 0, len(networks)) + for _, n := range networks { + names = append(names, n.Name) + } + + return names, nil +} diff --git a/management/server/store/sql_store_network_router_test.go b/management/server/store/sql_store_network_router_test.go new file mode 100644 index 000000000..c4e5280e6 --- /dev/null +++ b/management/server/store/sql_store_network_router_test.go @@ -0,0 +1,161 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetNetworkRoutersByNetID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + networkID string + expectedCount int + }{ + { + name: "retrieve routers by existing network ID", + networkID: "ct286bi7qv930dsrrug0", + expectedCount: 1, + }, + { + name: "retrieve routers by non-existing network ID", + networkID: "non-existent", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + routers, err := store.GetNetworkRoutersByNetID(context.Background(), LockingStrengthNone, accountID, tt.networkID) + require.NoError(t, err) + require.Len(t, routers, tt.expectedCount) + }) + } +} + +func TestSqlStore_GetNetworkRouterByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + networkRouterID string + expectError bool + }{ + { + name: "retrieve existing network router ID", + networkRouterID: "ctc20ji7qv9ck2sebc80", + expectError: false, + }, + { + name: "retrieve non-existing network router ID", + networkRouterID: "non-existing", + expectError: true, + }, + { + name: "retrieve network with empty router ID", + networkRouterID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + networkRouter, err := store.GetNetworkRouterByID(context.Background(), LockingStrengthNone, accountID, tt.networkRouterID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, networkRouter) + } else { + require.NoError(t, err) + require.NotNil(t, networkRouter) + require.Equal(t, tt.networkRouterID, networkRouter.ID) + } + }) + } +} + +func TestSqlStore_CreateNetworkRouter(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + networkID := "ct286bi7qv930dsrrug0" + + netRouter, err := routerTypes.NewNetworkRouter(accountID, networkID, "", []string{"net-router-grp"}, true, 0, true) + require.NoError(t, err) + + err = store.CreateNetworkRouter(context.Background(), netRouter) + require.NoError(t, err) + + savedNetRouter, err := store.GetNetworkRouterByID(context.Background(), LockingStrengthNone, accountID, netRouter.ID) + require.NoError(t, err) + require.Equal(t, netRouter, savedNetRouter) +} + +func TestSqlStore_UpdateNetworkRouter(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + networkID := "ct286bi7qv930dsrrug0" + routerID := "ctc20ji7qv9ck2sebc80" + + netRouter := &routerTypes.NetworkRouter{ + ID: routerID, + AccountID: accountID, + NetworkID: networkID, + Peer: "", + PeerGroups: []string{"net-router-grp"}, + Masquerade: true, + Metric: 42, + Enabled: true, + } + + err = store.UpdateNetworkRouter(context.Background(), netRouter) + require.NoError(t, err) + + savedNetRouter, err := store.GetNetworkRouterByID(context.Background(), LockingStrengthNone, accountID, routerID) + require.NoError(t, err) + require.Equal(t, netRouter, savedNetRouter) + + // Updating a router under a different account must not match any row. + netRouter.AccountID = "non-existent-account" + err = store.UpdateNetworkRouter(context.Background(), netRouter) + require.Error(t, err) +} + +func TestSqlStore_DeleteNetworkRouter(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + netRouterID := "ctc20ji7qv9ck2sebc80" + + err = store.DeleteNetworkRouter(context.Background(), accountID, netRouterID) + require.NoError(t, err) + + netRouter, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, netRouterID) + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, sErr.Type()) + require.Nil(t, netRouter) +} diff --git a/management/server/store/sql_store_network_test.go b/management/server/store/sql_store_network_test.go new file mode 100644 index 000000000..1eeda7cfc --- /dev/null +++ b/management/server/store/sql_store_network_test.go @@ -0,0 +1,128 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + networkTypes "github.com/netbirdio/netbird/management/server/networks/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetAccountNetworks(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectedCount int + }{ + { + name: "retrieve networks by existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectedCount: 1, + }, + + { + name: "retrieve networks by non-existing account ID", + accountID: "non-existent", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + networks, err := store.GetAccountNetworks(context.Background(), LockingStrengthNone, tt.accountID) + require.NoError(t, err) + require.Len(t, networks, tt.expectedCount) + }) + } +} + +func TestSqlStore_GetNetworkByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + networkID string + expectError bool + }{ + { + name: "retrieve existing network ID", + networkID: "ct286bi7qv930dsrrug0", + expectError: false, + }, + { + name: "retrieve non-existing network ID", + networkID: "non-existing", + expectError: true, + }, + { + name: "retrieve network with empty ID", + networkID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + network, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, tt.networkID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, network) + } else { + require.NoError(t, err) + require.NotNil(t, network) + require.Equal(t, tt.networkID, network.ID) + } + }) + } +} + +func TestSqlStore_SaveNetwork(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + network := &networkTypes.Network{ + ID: "net-id", + AccountID: accountID, + Name: "net", + } + + err = store.SaveNetwork(context.Background(), network) + require.NoError(t, err) + + savedNet, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, network.ID) + require.NoError(t, err) + require.Equal(t, network, savedNet) +} + +func TestSqlStore_DeleteNetwork(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + networkID := "ct286bi7qv930dsrrug0" + + err = store.DeleteNetwork(context.Background(), accountID, networkID) + require.NoError(t, err) + + network, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, networkID) + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, sErr.Type()) + require.Nil(t, network) +} diff --git a/management/server/store/sql_store_peer.go b/management/server/store/sql_store_peer.go new file mode 100644 index 000000000..9950069fd --- /dev/null +++ b/management/server/store/sql_store_peer.go @@ -0,0 +1,789 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "net" + "net/netip" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) SavePeer(ctx context.Context, accountID string, peer *nbpeer.Peer) error { + // To maintain data integrity, we create a copy of the peer's to prevent unintended updates to other fields. + peerCopy := peer.Copy() + peerCopy.AccountID = accountID + + err := s.transaction(ctx, func(tx *gorm.DB) error { + // check if peer exists before saving + var peerID string + result := tx.Model(&nbpeer.Peer{}).Select("id").Take(&peerID, accountAndIDQueryCondition, accountID, peer.ID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return status.Errorf(status.NotFound, peerNotFoundFMT, peer.ID) + } + return result.Error + } + + if peerID == "" { + return status.Errorf(status.NotFound, peerNotFoundFMT, peer.ID) + } + + result = tx.Model(&nbpeer.Peer{}).Where(accountAndIDQueryCondition, accountID, peer.ID).Save(peerCopy) + if result.Error != nil { + return status.Errorf(status.Internal, "failed to save peer to store: %v", result.Error) + } + + return nil + }) + if err != nil { + return err + } + + return nil +} + +func (s *SqlStore) SavePeerStatus(ctx context.Context, accountID, peerID string, peerStatus nbpeer.PeerStatus) error { + var peerCopy nbpeer.Peer + peerCopy.Status = &peerStatus + + fieldsToUpdate := []string{ + "peer_status_last_seen", "peer_status_session_started_at", + "peer_status_connected", "peer_status_login_expired", + "peer_status_requires_approval", + } + result := s.db.Model(&nbpeer.Peer{}). + Select(fieldsToUpdate). + Where(accountAndIDQueryCondition, accountID, peerID). + Updates(&peerCopy) + if result.Error != nil { + return status.Errorf(status.Internal, "failed to save peer status to store: %v", result.Error) + } + + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, peerNotFoundFMT, peerID) + } + + return nil +} + +// MarkPeerConnectedIfNewerSession is an atomic optimistic-locked update. +// The peer is marked connected with the given session token only when +// the stored SessionStartedAt is strictly smaller than the incoming +// one — equivalently, when no newer stream has already taken ownership. +// The sentinel zero (set on peer creation or after a disconnect) counts +// as the smallest possible token. This is the write half of the +// fencing protocol described on PeerStatus.SessionStartedAt. +// +// The post-write side effects in the caller — geo lookup, +// schedulePeerLoginExpiration, checkAndSchedulePeerInactivityExpiration, +// OnPeersUpdated — all run AFTER this method returns and are deliberately +// outside the database write so they cannot extend the row-lock window. +// +// LastSeen is set to the database's clock (CURRENT_TIMESTAMP) at the +// moment the row is written. The caller never supplies LastSeen because +// the value would otherwise drift under lock contention — a Go-side +// time.Now() taken before the write can land minutes later than the +// actual UPDATE under load, which previously caused real ordering bugs. +func (s *SqlStore) MarkPeerConnectedIfNewerSession(ctx context.Context, accountID, peerID string, newSessionStartedAt int64) (bool, error) { + result := s.db.WithContext(ctx). + Model(&nbpeer.Peer{}). + Where(accountAndIDQueryCondition, accountID, peerID). + Where("peer_status_session_started_at < ?", newSessionStartedAt). + Updates(map[string]any{ + "peer_status_connected": true, + "peer_status_last_seen": gorm.Expr("CURRENT_TIMESTAMP"), + "peer_status_session_started_at": newSessionStartedAt, + "peer_status_login_expired": false, + }) + if result.Error != nil { + return false, status.Errorf(status.Internal, "mark peer connected: %v", result.Error) + } + return result.RowsAffected > 0, nil +} + +// MarkPeerDisconnectedIfSameSession is an atomic optimistic-locked update. +// The peer is marked disconnected only when the stored SessionStartedAt +// matches the incoming token — meaning the stream that owns the current +// session is the one ending. If a newer stream has already replaced the +// session, the update is skipped. LastSeen is set to CURRENT_TIMESTAMP at +// write time; see MarkPeerConnectedIfNewerSession for the rationale. +// +// A zero sessionStartedAt is rejected at the call site; the underlying +// WHERE on equality would otherwise match every never-connected peer. +func (s *SqlStore) MarkPeerDisconnectedIfSameSession(ctx context.Context, accountID, peerID string, sessionStartedAt int64) (bool, error) { + if sessionStartedAt == 0 { + return false, nil + } + result := s.db.WithContext(ctx). + Model(&nbpeer.Peer{}). + Where(accountAndIDQueryCondition, accountID, peerID). + Where("peer_status_session_started_at = ?", sessionStartedAt). + Updates(map[string]any{ + "peer_status_connected": false, + "peer_status_last_seen": gorm.Expr("CURRENT_TIMESTAMP"), + "peer_status_session_started_at": int64(0), + }) + if result.Error != nil { + return false, status.Errorf(status.Internal, "mark peer disconnected: %v", result.Error) + } + return result.RowsAffected > 0, nil +} + +// ApproveAccountPeers marks all peers that currently require approval in the given account as approved. +func (s *SqlStore) ApproveAccountPeers(ctx context.Context, accountID string) (int, error) { + result := s.db.Model(&nbpeer.Peer{}). + Where("account_id = ? AND peer_status_requires_approval = ?", accountID, true). + Update("peer_status_requires_approval", false) + if result.Error != nil { + return 0, status.Errorf(status.Internal, "failed to approve pending account peers: %v", result.Error) + } + + return int(result.RowsAffected), nil +} + +// RefreshPeerLastSeen updates only peer_status_last_seen. Every other status +// column is left untouched: peer_status_connected and +// peer_status_session_started_at belong to the sync stream that owns the +// session, and a blind write here would corrupt the fencing +// MarkPeerConnectedIfNewerSession relies on. +// +// LastSeen comes from the database clock for the same reason it does there: a +// Go-side timestamp is taken before the write and can land after a connect that +// used CURRENT_TIMESTAMP, dragging the column backwards. +// +// staleBefore carries the caller's throttle into the same statement, so +// concurrent requests for one peer collapse into a single write instead of +// each racing on its own stale read. The column is nullable — Status is an +// embedded pointer, so a peer stored without one leaves it NULL — and NULL +// loses every comparison, hence the explicit branch for a peer never seen. +func (s *SqlStore) RefreshPeerLastSeen(ctx context.Context, accountID, peerID string, staleBefore time.Time) (bool, error) { + result := s.db.WithContext(ctx). + Model(&nbpeer.Peer{}). + Where(accountAndIDQueryCondition, accountID, peerID). + Where("(peer_status_last_seen IS NULL OR peer_status_last_seen < ?)", staleBefore). + Update("peer_status_last_seen", gorm.Expr("CURRENT_TIMESTAMP")) + if result.Error != nil { + return false, status.Errorf(status.Internal, "refresh peer last seen: %v", result.Error) + } + + return result.RowsAffected > 0, nil +} + +func (s *SqlStore) getPeers(ctx context.Context, accountID string) ([]nbpeer.Peer, error) { + const query = `SELECT id, account_id, key, ip, name, dns_label, user_id, ssh_key, ssh_enabled, login_expiration_enabled, + inactivity_expiration_enabled, last_login, created_at, ephemeral, extra_dns_labels, allow_extra_dns_labels, meta_hostname, + meta_go_os, meta_kernel, meta_core, meta_platform, meta_os, meta_os_version, meta_wt_version, meta_ui_version, + meta_kernel_version, meta_network_addresses, meta_system_serial_number, meta_system_product_name, meta_system_manufacturer, + meta_environment, meta_flags, meta_files, meta_certificates, meta_capabilities, peer_status_last_seen, peer_status_session_started_at, + peer_status_connected, peer_status_login_expired, peer_status_requires_approval, location_connection_ip, + location_country_code, location_city_name, location_geo_name_id, proxy_meta_embedded, proxy_meta_cluster, ipv6, meta_sync_message_version + FROM peers WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + + peers, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nbpeer.Peer, error) { + var p nbpeer.Peer + p.Status = &nbpeer.PeerStatus{} + var ( + lastLogin, createdAt sql.NullTime + sshEnabled, loginExpirationEnabled, inactivityExpirationEnabled, ephemeral, allowExtraDNSLabels sql.NullBool + peerStatusLastSeen sql.NullTime + peerStatusSessionStartedAt sql.NullInt64 + peerStatusConnected, peerStatusLoginExpired, peerStatusRequiresApproval, proxyEmbedded sql.NullBool + ip, extraDNS, netAddr, env, flags, files, certificates, capabilities, connIP, ipv6 []byte + metaHostname, metaGoOS, metaKernel, metaCore, metaPlatform sql.NullString + metaOS, metaOSVersion, metaWtVersion, metaUIVersion, metaKernelVersion sql.NullString + metaSystemSerialNumber, metaSystemProductName, metaSystemManufacturer sql.NullString + locationCountryCode, locationCityName, proxyCluster sql.NullString + locationGeoNameID sql.NullInt64 + metaSyncMessageVersion sql.NullInt32 + ) + + err := row.Scan(&p.ID, &p.AccountID, &p.Key, &ip, &p.Name, &p.DNSLabel, &p.UserID, &p.SSHKey, &sshEnabled, + &loginExpirationEnabled, &inactivityExpirationEnabled, &lastLogin, &createdAt, &ephemeral, &extraDNS, + &allowExtraDNSLabels, &metaHostname, &metaGoOS, &metaKernel, &metaCore, &metaPlatform, + &metaOS, &metaOSVersion, &metaWtVersion, &metaUIVersion, &metaKernelVersion, &netAddr, + &metaSystemSerialNumber, &metaSystemProductName, &metaSystemManufacturer, &env, &flags, &files, &certificates, &capabilities, + &peerStatusLastSeen, &peerStatusSessionStartedAt, &peerStatusConnected, &peerStatusLoginExpired, + &peerStatusRequiresApproval, &connIP, &locationCountryCode, &locationCityName, &locationGeoNameID, + &proxyEmbedded, &proxyCluster, &ipv6, &metaSyncMessageVersion) + + if err == nil { + if lastLogin.Valid { + p.LastLogin = &lastLogin.Time + } + if createdAt.Valid { + p.CreatedAt = createdAt.Time + } + if sshEnabled.Valid { + p.SSHEnabled = sshEnabled.Bool + } + if loginExpirationEnabled.Valid { + p.LoginExpirationEnabled = loginExpirationEnabled.Bool + } + if inactivityExpirationEnabled.Valid { + p.InactivityExpirationEnabled = inactivityExpirationEnabled.Bool + } + if ephemeral.Valid { + p.Ephemeral = ephemeral.Bool + } + if allowExtraDNSLabels.Valid { + p.AllowExtraDNSLabels = allowExtraDNSLabels.Bool + } + if peerStatusLastSeen.Valid { + p.Status.LastSeen = peerStatusLastSeen.Time + } + if peerStatusSessionStartedAt.Valid { + p.Status.SessionStartedAt = peerStatusSessionStartedAt.Int64 + } + if peerStatusConnected.Valid { + p.Status.Connected = peerStatusConnected.Bool + } + if peerStatusLoginExpired.Valid { + p.Status.LoginExpired = peerStatusLoginExpired.Bool + } + if peerStatusRequiresApproval.Valid { + p.Status.RequiresApproval = peerStatusRequiresApproval.Bool + } + if metaHostname.Valid { + p.Meta.Hostname = metaHostname.String + } + if metaGoOS.Valid { + p.Meta.GoOS = metaGoOS.String + } + if metaKernel.Valid { + p.Meta.Kernel = metaKernel.String + } + if metaCore.Valid { + p.Meta.Core = metaCore.String + } + if metaPlatform.Valid { + p.Meta.Platform = metaPlatform.String + } + if metaOS.Valid { + p.Meta.OS = metaOS.String + } + if metaOSVersion.Valid { + p.Meta.OSVersion = metaOSVersion.String + } + if metaWtVersion.Valid { + p.Meta.WtVersion = metaWtVersion.String + } + if metaUIVersion.Valid { + p.Meta.UIVersion = metaUIVersion.String + } + if metaKernelVersion.Valid { + p.Meta.KernelVersion = metaKernelVersion.String + } + if metaSystemSerialNumber.Valid { + p.Meta.SystemSerialNumber = metaSystemSerialNumber.String + } + if metaSystemProductName.Valid { + p.Meta.SystemProductName = metaSystemProductName.String + } + if metaSystemManufacturer.Valid { + p.Meta.SystemManufacturer = metaSystemManufacturer.String + } + if locationCountryCode.Valid { + p.Location.CountryCode = locationCountryCode.String + } + if locationCityName.Valid { + p.Location.CityName = locationCityName.String + } + if locationGeoNameID.Valid { + p.Location.GeoNameID = uint(locationGeoNameID.Int64) + } + if proxyEmbedded.Valid { + p.ProxyMeta.Embedded = proxyEmbedded.Bool + } + if proxyCluster.Valid { + p.ProxyMeta.Cluster = proxyCluster.String + } + if ip != nil { + _ = json.Unmarshal(ip, &p.IP) + } + if ipv6 != nil { + _ = json.Unmarshal(ipv6, &p.IPv6) + } + if extraDNS != nil { + _ = json.Unmarshal(extraDNS, &p.ExtraDNSLabels) + } + if netAddr != nil { + _ = json.Unmarshal(netAddr, &p.Meta.NetworkAddresses) + } + if env != nil { + _ = json.Unmarshal(env, &p.Meta.Environment) + } + if flags != nil { + _ = json.Unmarshal(flags, &p.Meta.Flags) + } + if files != nil { + _ = json.Unmarshal(files, &p.Meta.Files) + } + if certificates != nil { + _ = json.Unmarshal(certificates, &p.Meta.Certificates) + } + if capabilities != nil { + _ = json.Unmarshal(capabilities, &p.Meta.Capabilities) + } + if connIP != nil { + _ = json.Unmarshal(connIP, &p.Location.ConnectionIP) + } + if metaSyncMessageVersion.Valid { + p.Meta.SyncMessageVersion = int(metaSyncMessageVersion.Int32) + } + } + return p, err + }) + if err != nil { + return nil, err + } + return peers, nil +} + +func (s *SqlStore) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) { + var peer nbpeer.Peer + result := s.db.Select("account_id").Take(&peer, idQueryCondition, peerID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + return nil, status.NewGetAccountFromStoreError(result.Error) + } + + if peer.AccountID == "" { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + + return s.GetAccount(ctx, peer.AccountID) +} + +func (s *SqlStore) GetAccountByPeerPubKey(ctx context.Context, peerKey string) (*types.Account, error) { + var peer nbpeer.Peer + result := s.db.Select("account_id").Take(&peer, GetKeyQueryCondition(s), peerKey) + + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + return nil, status.NewGetAccountFromStoreError(result.Error) + } + + if peer.AccountID == "" { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + + return s.GetAccount(ctx, peer.AccountID) +} + +func (s *SqlStore) GetAccountIDByPeerPubKey(ctx context.Context, peerKey string) (string, error) { + var peer nbpeer.Peer + var accountID string + result := s.db.Model(&peer).Select("account_id").Where(GetKeyQueryCondition(s), peerKey).Take(&accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.Errorf(status.NotFound, "account not found: index lookup failed") + } + return "", status.NewGetAccountFromStoreError(result.Error) + } + + return accountID, nil +} + +func (s *SqlStore) GetAccountIDByPeerID(ctx context.Context, lockStrength LockingStrength, peerID string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountID string + result := tx.Model(&nbpeer.Peer{}). + Select("account_id").Where(idQueryCondition, peerID).Take(&accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.Errorf(status.NotFound, "peer %s account not found", peerID) + } + return "", status.NewGetAccountFromStoreError(result.Error) + } + + return accountID, nil +} + +func (s *SqlStore) GetTakenIPs(ctx context.Context, lockStrength LockingStrength, accountID string) ([]netip.Addr, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var ipJSONStrings []string + + result := tx.Model(&nbpeer.Peer{}). + Where("account_id = ?", accountID). + Pluck("ip", &ipJSONStrings) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "no peers found for the account") + } + return nil, status.Errorf(status.Internal, "issue getting IPs from store: %s", result.Error) + } + + ips := make([]netip.Addr, len(ipJSONStrings)) + for i, ipJSON := range ipJSONStrings { + var ip netip.Addr + if err := json.Unmarshal([]byte(ipJSON), &ip); err != nil { + return nil, status.Errorf(status.Internal, "issue parsing IP JSON from store") + } + ips[i] = ip.Unmap() + } + + return ips, nil +} + +func (s *SqlStore) GetPeerLabelsInAccount(ctx context.Context, lockStrength LockingStrength, accountID string, dnsLabel string) ([]string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var labels []string + result := tx.Model(&nbpeer.Peer{}). + Where("account_id = ? AND dns_label LIKE ?", accountID, dnsLabel+"%"). + Pluck("dns_label", &labels) + + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "no peers found for the account") + } + log.WithContext(ctx).Errorf("error when getting dns labels from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "issue getting dns labels from store: %s", result.Error) + } + + return labels, nil +} + +func (s *SqlStore) GetPeerByPeerPubKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peer nbpeer.Peer + result := tx.Take(&peer, GetKeyQueryCondition(s), peerKey) + + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewPeerNotFoundError(peerKey) + } + return nil, status.Errorf(status.Internal, "issue getting peer from store: %s", result.Error) + } + + return &peer, nil +} + +// GetAccountPeers retrieves peers for an account. +func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { + var peers []*nbpeer.Peer + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + query := tx.Where(accountIDCondition, accountID) + + if nameFilter != "" { + query = query.Where("name LIKE ?", "%"+nameFilter+"%") + } + if ipFilter != "" { + query = query.Where("ip LIKE ? OR ipv6 LIKE ?", "%"+ipFilter+"%", "%"+ipFilter+"%") + } + // MAC addresses live in the JSON-serialized meta_network_addresses column, + // so we match the raw JSON text rather than a dedicated column. + if macFilter != "" { + query = query.Where("meta_network_addresses LIKE ?", "%"+macFilter+"%") + } + + if err := query.Find(&peers).Error; err != nil { + log.WithContext(ctx).Errorf("failed to get peers from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get peers from store") + } + + return peers, nil +} + +// GetUserPeers retrieves peers for a user. +func (s *SqlStore) GetUserPeers(ctx context.Context, lockStrength LockingStrength, accountID, userID string) ([]*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peers []*nbpeer.Peer + + // Exclude peers added via setup keys, as they are not user-specific and have an empty user_id. + if userID == "" { + return peers, nil + } + + result := tx. + Find(&peers, "account_id = ? AND user_id = ?", accountID, userID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get peers from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get peers from store") + } + + return peers, nil +} + +func (s *SqlStore) AddPeerToAccount(ctx context.Context, peer *nbpeer.Peer) error { + if err := s.db.Create(peer).Error; err != nil { + return status.Errorf(status.Internal, "issue adding peer to account: %s", err) + } + + return nil +} + +// GetPeerByID retrieves a peer by its ID and account ID. +func (s *SqlStore) GetPeerByID(ctx context.Context, lockStrength LockingStrength, accountID, peerID string) (*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peer *nbpeer.Peer + result := tx. + Take(&peer, accountAndIDQueryCondition, accountID, peerID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewPeerNotFoundError(peerID) + } + return nil, status.Errorf(status.Internal, "failed to get peer from store") + } + + return peer, nil +} + +// GetPeersByIDs retrieves peers by their IDs and account ID. +func (s *SqlStore) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peers []*nbpeer.Peer + result := tx.Find(&peers, accountAndIDsQueryCondition, accountID, peerIDs) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get peers by ID's from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get peers by ID's from the store") + } + + peersMap := make(map[string]*nbpeer.Peer) + for _, peer := range peers { + peersMap[peer.ID] = peer + } + + return peersMap, nil +} + +// GetAccountPeersWithExpiration retrieves a list of peers that have login expiration enabled and added by a user. +func (s *SqlStore) GetAccountPeersWithExpiration(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peers []*nbpeer.Peer + result := tx. + Where("login_expiration_enabled = ? AND peer_status_login_expired != ? AND user_id IS NOT NULL AND user_id != ''", true, true). + Find(&peers, accountIDCondition, accountID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get peers with expiration from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get peers with expiration from store") + } + + return peers, nil +} + +// GetAccountPeersWithInactivity retrieves a list of peers that have login expiration enabled and added by a user. +func (s *SqlStore) GetAccountPeersWithInactivity(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peers []*nbpeer.Peer + result := tx. + Where("inactivity_expiration_enabled = ? AND user_id IS NOT NULL AND user_id != ''", true). + Find(&peers, accountIDCondition, accountID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get peers with inactivity from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get peers with inactivity from store") + } + + return peers, nil +} + +// GetAllEphemeralPeers retrieves all peers with Ephemeral set to true across all accounts, optimized for batch processing. +func (s *SqlStore) GetAllEphemeralPeers(ctx context.Context, lockStrength LockingStrength) ([]*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var allEphemeralPeers, batchPeers []*nbpeer.Peer + result := tx. + Where("ephemeral = ?", true). + FindInBatches(&batchPeers, 1000, func(tx *gorm.DB, batch int) error { + allEphemeralPeers = append(allEphemeralPeers, batchPeers...) + return nil + }) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to retrieve ephemeral peers: %s", result.Error) + return nil, fmt.Errorf("failed to retrieve ephemeral peers") + } + + return allEphemeralPeers, nil +} + +// DeletePeer removes a peer from the store. +func (s *SqlStore) DeletePeer(ctx context.Context, accountID string, peerID string) error { + result := s.db.Delete(&nbpeer.Peer{}, accountAndIDQueryCondition, accountID, peerID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to delete peer from the store: %s", err) + return status.Errorf(status.Internal, "failed to delete peer from store") + } + + if result.RowsAffected == 0 { + return status.NewPeerNotFoundError(peerID) + } + + return nil +} + +func (s *SqlStore) GetPeerByIP(ctx context.Context, lockStrength LockingStrength, accountID string, ip net.IP) (*nbpeer.Peer, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + column := "ip" + if ip.To4() == nil { + column = "ipv6" + } + jsonValue := fmt.Sprintf(`"%s"`, ip.String()) + + var peer nbpeer.Peer + result := tx. + Take(&peer, fmt.Sprintf("account_id = ? AND %s = ?", column), accountID, jsonValue) + if result.Error != nil { + // A tunnel-IP miss is an expected outcome (e.g. the proxy's + // ValidateTunnelPeer probing an address that isn't in the + // account roster); surface it as NotFound so callers can tell + // it apart from a real store failure. + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "peer with ip %s not found", ip.String()) + } + return nil, status.Errorf(status.Internal, "failed to get peer from store") + } + + return &peer, nil +} + +func (s *SqlStore) GetPeerIdByLabel(ctx context.Context, lockStrength LockingStrength, accountID string, hostname string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peerID string + result := tx.Model(&nbpeer.Peer{}). + Select("id"). + // Where(" = ?", hostname). + Where("account_id = ? AND dns_label = ?", accountID, hostname). + Limit(1). + Scan(&peerID) + + if peerID == "" { + return "", gorm.ErrRecordNotFound + } + + return peerID, result.Error +} + +// GetEmbeddedProxyPeerIDsByCluster returns peer IDs of all embedded proxy peers +// in the account, grouped by their ProxyCluster. The map is nil when no embedded +// proxy peers exist. +func (s *SqlStore) GetEmbeddedProxyPeerIDsByCluster(ctx context.Context, accountID string) (map[string][]string, error) { + type row struct { + ID string + Cluster string + } + var rows []row + result := s.db.Model(&nbpeer.Peer{}). + Select("id, proxy_meta_cluster AS cluster"). + Where("account_id = ? AND proxy_meta_embedded = ?", accountID, true). + Scan(&rows) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "failed to get embedded proxy peers: %s", result.Error) + } + + out := make(map[string][]string, len(rows)) + for _, r := range rows { + out[r.Cluster] = append(out[r.Cluster], r.ID) + } + return out, nil +} + +func (s *SqlStore) GetUserIDByPeerKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var userID string + result := tx.Model(&nbpeer.Peer{}). + Select("user_id"). + Take(&userID, GetKeyQueryCondition(s), peerKey) + + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.Errorf(status.NotFound, "peer not found: index lookup failed") + } + return "", status.Errorf(status.Internal, "failed to get user ID by peer key") + } + + return userID, nil +} + +func (s *SqlStore) GetPeerIDByKey(ctx context.Context, lockStrength LockingStrength, key string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var peerID string + result := tx.Model(&nbpeer.Peer{}). + Select("id"). + Where(GetKeyQueryCondition(s), key). + Limit(1). + Scan(&peerID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get peer ID by key: %s", result.Error) + return "", status.Errorf(status.Internal, "failed to get peer ID by key") + } + + return peerID, nil +} diff --git a/management/server/store/sql_store_peer_test.go b/management/server/store/sql_store_peer_test.go new file mode 100644 index 000000000..1432b5d96 --- /dev/null +++ b/management/server/store/sql_store_peer_test.go @@ -0,0 +1,943 @@ +package store + +import ( + "context" + "encoding/binary" + "fmt" + "net" + "net/netip" + "reflect" + "sort" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/management/server/util" + "github.com/netbirdio/netbird/shared/management/status" + "github.com/netbirdio/netbird/shared/testing_helpers" +) + +// TestSqlStore_GetPeerByIP_NotFound pins the not-found semantics the +// proxy's ValidateTunnelPeer relies on: a tunnel-IP that isn't in the +// account roster must surface as a NotFound error (not a generic +// Internal) so callers can distinguish an expected miss from a real +// store failure. A known IP still resolves. +func TestSqlStore_GetPeerByIP_NotFound(t *testing.T) { + runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { + const accountID = "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + peer, err := store.GetPeerByIP(context.Background(), LockingStrengthNone, accountID, net.ParseIP("192.168.0.0")) + require.NoError(t, err, "known tunnel IP must resolve") + require.NotNil(t, peer) + + _, err = store.GetPeerByIP(context.Background(), LockingStrengthNone, accountID, net.ParseIP("100.65.0.99")) + require.Error(t, err, "unknown tunnel IP must error") + parsedErr, ok := status.FromError(err) + require.True(t, ok, "error must be a status error") + require.Equal(t, status.NotFound, parsedErr.Type(), "tunnel-IP miss must be NotFound, not Internal") + }) +} + +func TestSqlStore_SavePeer(t *testing.T) { + populateFields := testing_helpers.NewPopulateFields() + + runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { + account, err := store.GetAccount(context.Background(), "bf1c8084-ba50-4ce7-9439-34653001fc3b") + require.NoError(t, err) + + metadata := nbpeer.PeerSystemMeta{} + reflectedMetadata := reflect.ValueOf(&metadata).Elem() + + numOfFields, err := populateFields.PopulateAll(reflectedMetadata) + assert.NoError(t, err) + assert.Equal(t, 33, numOfFields) + + // save status of non-existing peer + peer := &nbpeer.Peer{ + Key: "peerkey", + ID: "testpeer", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::1"), + Meta: metadata, //nbpeer.PeerSystemMeta{Hostname: "testingpeer"}, + Name: "peer name", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + CreatedAt: time.Now().UTC(), + } + ctx := context.Background() + err = store.SavePeer(ctx, account.Id, peer) + assert.Error(t, err) + parsedErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") + + // save new status of existing peer + account.Peers[peer.ID] = peer + + err = store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + updatedPeer := peer.Copy() + updatedPeer.Status.Connected = false + updatedPeer.Meta.Hostname = "updatedpeer" + + err = store.SavePeer(ctx, account.Id, updatedPeer) + require.NoError(t, err) + + account, err = store.GetAccount(context.Background(), account.Id) + require.NoError(t, err) + + actual := account.Peers[peer.ID] + assert.Equal(t, updatedPeer.Meta, actual.Meta) + assert.Equal(t, updatedPeer.Status.Connected, actual.Status.Connected) + assert.Equal(t, updatedPeer.Status.LoginExpired, actual.Status.LoginExpired) + assert.Equal(t, updatedPeer.Status.RequiresApproval, actual.Status.RequiresApproval) + assert.WithinDurationf(t, updatedPeer.Status.LastSeen, actual.Status.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") + }) +} + +func TestSqlStore_SavePeerStatus(t *testing.T) { + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + account, err := store.GetAccount(context.Background(), "bf1c8084-ba50-4ce7-9439-34653001fc3b") + require.NoError(t, err) + + // save status of non-existing peer + newStatus := nbpeer.PeerStatus{Connected: false, LastSeen: time.Now().UTC()} + err = store.SavePeerStatus(context.Background(), account.Id, "non-existing-peer", newStatus) + assert.Error(t, err) + parsedErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") + + // save new status of existing peer + account.Peers["testpeer"] = &nbpeer.Peer{ + Key: "peerkey", + ID: "testpeer", + IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::1"), + Meta: nbpeer.PeerSystemMeta{}, + Name: "peer name", + Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, + } + + err = store.SaveAccount(context.Background(), account) + require.NoError(t, err) + + err = store.SavePeerStatus(context.Background(), account.Id, "testpeer", newStatus) + require.NoError(t, err) + + account, err = store.GetAccount(context.Background(), account.Id) + require.NoError(t, err) + + actual := account.Peers["testpeer"].Status + assert.Equal(t, newStatus.Connected, actual.Connected) + assert.Equal(t, newStatus.LoginExpired, actual.LoginExpired) + assert.Equal(t, newStatus.RequiresApproval, actual.RequiresApproval) + assert.WithinDurationf(t, newStatus.LastSeen, actual.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") + + newStatus.Connected = true + + err = store.SavePeerStatus(context.Background(), account.Id, "testpeer", newStatus) + require.NoError(t, err) + + account, err = store.GetAccount(context.Background(), account.Id) + require.NoError(t, err) + + actual = account.Peers["testpeer"].Status + assert.Equal(t, newStatus.Connected, actual.Connected) + assert.Equal(t, newStatus.LoginExpired, actual.LoginExpired) + assert.Equal(t, newStatus.RequiresApproval, actual.RequiresApproval) + assert.WithinDurationf(t, newStatus.LastSeen, actual.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") +} + +func TestSqlite_GetTakenIPs(t *testing.T) { + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + defer cleanup() + if err != nil { + t.Fatal(err) + } + + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + _, err = store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + takenIPs, err := store.GetTakenIPs(context.Background(), LockingStrengthNone, existingAccountID) + require.NoError(t, err) + assert.Equal(t, []netip.Addr{}, takenIPs) + + peer1 := &nbpeer.Peer{ + ID: "peer1", + AccountID: existingAccountID, + Key: "key1", + DNSLabel: "peer1", + IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), + IPv6: netip.MustParseAddr("fd00::1:1:1:1"), + } + err = store.AddPeerToAccount(context.Background(), peer1) + require.NoError(t, err) + + takenIPs, err = store.GetTakenIPs(context.Background(), LockingStrengthNone, existingAccountID) + require.NoError(t, err) + ip1 := netip.AddrFrom4([4]byte{1, 1, 1, 1}) + assert.Equal(t, []netip.Addr{ip1}, takenIPs) + + peer2 := &nbpeer.Peer{ + ID: "peer1second", + AccountID: existingAccountID, + Key: "key2", + DNSLabel: "peer1-1", + IP: netip.AddrFrom4([4]byte{2, 2, 2, 2}), + IPv6: netip.MustParseAddr("fd00::2:2:2:2"), + } + err = store.AddPeerToAccount(context.Background(), peer2) + require.NoError(t, err) + + takenIPs, err = store.GetTakenIPs(context.Background(), LockingStrengthNone, existingAccountID) + require.NoError(t, err) + ip2 := netip.AddrFrom4([4]byte{2, 2, 2, 2}) + assert.Equal(t, []netip.Addr{ip1, ip2}, takenIPs) +} + +func TestSqlite_GetPeerLabelsInAccount(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + peerHostname := "peer1" + + _, err := store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + labels, err := store.GetPeerLabelsInAccount(context.Background(), LockingStrengthNone, existingAccountID, peerHostname) + require.NoError(t, err) + assert.Equal(t, []string{}, labels) + + peer1 := &nbpeer.Peer{ + ID: "peer1", + AccountID: existingAccountID, + Key: "key1", + DNSLabel: "peer1", + IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), + IPv6: netip.MustParseAddr("fd00::1:1:1:1"), + } + err = store.AddPeerToAccount(context.Background(), peer1) + require.NoError(t, err) + + labels, err = store.GetPeerLabelsInAccount(context.Background(), LockingStrengthNone, existingAccountID, peerHostname) + require.NoError(t, err) + assert.Equal(t, []string{"peer1"}, labels) + + peer2 := &nbpeer.Peer{ + ID: "peer1second", + AccountID: existingAccountID, + Key: "key2", + DNSLabel: "peer1-1", + IP: netip.AddrFrom4([4]byte{2, 2, 2, 2}), + IPv6: netip.MustParseAddr("fd00::2:2:2:2"), + } + err = store.AddPeerToAccount(context.Background(), peer2) + require.NoError(t, err) + + labels, err = store.GetPeerLabelsInAccount(context.Background(), LockingStrengthNone, existingAccountID, peerHostname) + require.NoError(t, err) + + expected := []string{"peer1", "peer1-1"} + sort.Strings(expected) + sort.Strings(labels) + assert.Equal(t, expected, labels) + }) +} + +func Test_AddPeerWithSameDnsLabel(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + _, err := store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + peer1 := &nbpeer.Peer{ + ID: "peer1", + AccountID: existingAccountID, + Key: "key1", + DNSLabel: "peer1.domain.test", + } + err = store.AddPeerToAccount(context.Background(), peer1) + require.NoError(t, err) + + peer2 := &nbpeer.Peer{ + ID: "peer1second", + AccountID: existingAccountID, + Key: "key2", + DNSLabel: "peer1.domain.test", + } + err = store.AddPeerToAccount(context.Background(), peer2) + require.Error(t, err) + }) +} + +func Test_AddPeerWithSameIP(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + _, err := store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + peer1 := &nbpeer.Peer{ + ID: "peer1", + AccountID: existingAccountID, + Key: "key1", + IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), + IPv6: netip.MustParseAddr("fd00::1:1:1:1"), + } + err = store.AddPeerToAccount(context.Background(), peer1) + require.NoError(t, err) + + peer2 := &nbpeer.Peer{ + ID: "peer1second", + AccountID: existingAccountID, + Key: "key2", + IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), + IPv6: netip.MustParseAddr("fd00::2:2:2:2"), + } + err = store.AddPeerToAccount(context.Background(), peer2) + require.Error(t, err) + }) +} + +func TestSqlStore_GetPeerByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + peerID string + expectError bool + }{ + { + name: "retrieve existing peer", + peerID: "cfefqs706sqkneg59g4g", + expectError: false, + }, + { + name: "retrieve non-existing peer", + peerID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty peer ID", + peerID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peer, err := store.GetPeerByID(context.Background(), LockingStrengthNone, accountID, tt.peerID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, peer) + } else { + require.NoError(t, err) + require.NotNil(t, peer) + require.Equal(t, tt.peerID, peer.ID) + } + }) + } +} + +func TestSqlStore_GetPeersByIDs(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + peerIDs []string + expectedCount int + }{ + { + name: "retrieve existing peers by existing IDs", + peerIDs: []string{"cfefqs706sqkneg59g4g", "cfeg6sf06sqkneg59g50"}, + expectedCount: 2, + }, + { + name: "empty peer IDs list", + peerIDs: []string{}, + expectedCount: 0, + }, + { + name: "non-existing peer IDs", + peerIDs: []string{"nonexistent1", "nonexistent2"}, + expectedCount: 0, + }, + { + name: "mixed existing and non-existing peer IDs", + peerIDs: []string{"cfeg6sf06sqkneg59g50", "nonexistent"}, + expectedCount: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetPeersByIDs(context.Background(), LockingStrengthNone, accountID, tt.peerIDs) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } +} + +func TestSqlStore_AddPeerToAccount(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + peer := &nbpeer.Peer{ + ID: "peer1", + AccountID: accountID, + Key: "key", + IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), + IPv6: netip.MustParseAddr("fd00::1:1:1:1"), + Meta: nbpeer.PeerSystemMeta{ + Hostname: "hostname", + GoOS: "linux", + Kernel: "Linux", + Core: "21.04", + Platform: "x86_64", + OS: "Ubuntu", + WtVersion: "development", + UIVersion: "development", + }, + Name: "peer.test", + DNSLabel: "peer", + Status: &nbpeer.PeerStatus{ + LastSeen: time.Now().UTC(), + Connected: true, + LoginExpired: false, + RequiresApproval: false, + }, + SSHKey: "ssh-key", + SSHEnabled: false, + LoginExpirationEnabled: true, + InactivityExpirationEnabled: false, + LastLogin: util.ToPtr(time.Now().UTC()), + CreatedAt: time.Now().UTC(), + Ephemeral: true, + } + err = store.AddPeerToAccount(context.Background(), peer) + require.NoError(t, err, "failed to add peer to account") + + storedPeer, err := store.GetPeerByID(context.Background(), LockingStrengthNone, accountID, peer.ID) + require.NoError(t, err, "failed to get peer") + + assert.Equal(t, peer.ID, storedPeer.ID) + assert.Equal(t, peer.AccountID, storedPeer.AccountID) + assert.Equal(t, peer.Key, storedPeer.Key) + assert.Equal(t, peer.IP.String(), storedPeer.IP.String()) + assert.Equal(t, peer.Meta, storedPeer.Meta) + assert.Equal(t, peer.Name, storedPeer.Name) + assert.Equal(t, peer.DNSLabel, storedPeer.DNSLabel) + assert.Equal(t, peer.SSHKey, storedPeer.SSHKey) + assert.Equal(t, peer.SSHEnabled, storedPeer.SSHEnabled) + assert.Equal(t, peer.LoginExpirationEnabled, storedPeer.LoginExpirationEnabled) + assert.Equal(t, peer.InactivityExpirationEnabled, storedPeer.InactivityExpirationEnabled) + assert.WithinDurationf(t, peer.GetLastLogin(), storedPeer.GetLastLogin().UTC(), time.Millisecond, "LastLogin should be equal") + assert.WithinDurationf(t, peer.CreatedAt, storedPeer.CreatedAt.UTC(), time.Millisecond, "CreatedAt should be equal") + assert.Equal(t, peer.Ephemeral, storedPeer.Ephemeral) + assert.Equal(t, peer.Status.Connected, storedPeer.Status.Connected) + assert.Equal(t, peer.Status.LoginExpired, storedPeer.Status.LoginExpired) + assert.Equal(t, peer.Status.RequiresApproval, storedPeer.Status.RequiresApproval) + assert.WithinDurationf(t, peer.Status.LastSeen, storedPeer.Status.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") +} + +func TestSqlStore_GetAccountPeers(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + nameFilter string + ipFilter string + expectedCount int + }{ + { + name: "should retrieve peers for an existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectedCount: 5, + }, + { + name: "should return no peers for a non-existing account ID", + accountID: "nonexistent", + expectedCount: 0, + }, + { + name: "should return no peers for an empty account ID", + accountID: "", + expectedCount: 0, + }, + { + name: "should filter peers by name", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + nameFilter: "expiredhost", + expectedCount: 1, + }, + { + name: "should filter peers by partial name", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + nameFilter: "host", + expectedCount: 4, + }, + { + name: "should filter peers by ip", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + ipFilter: "100.64.39.54", + expectedCount: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetAccountPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.nameFilter, tt.ipFilter, "") + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } + +} + +func TestSqlStore_GetAccountPeers_FilterByMac(t *testing.T) { + ctx := context.Background() + store, cleanup, err := NewTestStoreFromSQL(ctx, "", t.TempDir()) + require.NoError(t, err) + t.Cleanup(cleanup) + + accountID := "test-account-mac" + userID := "test-user-mac" + account := newAccountWithId(ctx, accountID, userID, "example.com") + account.Peers["peer-mac-1"] = &nbpeer.Peer{ + ID: "peer-mac-1", + AccountID: accountID, + Key: "peer-mac-key-1", + Name: "macpeer", + IP: netip.MustParseAddr("100.64.0.10"), + Meta: nbpeer.PeerSystemMeta{ + NetworkAddresses: []nbpeer.NetworkAddress{ + {NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"}, + }, + }, + } + require.NoError(t, store.SaveAccount(ctx, account)) + + tests := []struct { + name string + macFilter string + expectedCount int + }{ + {name: "full mac matches", macFilter: "00:93:37:bd:83:0f", expectedCount: 1}, + {name: "mac prefix matches", macFilter: "00:93:37", expectedCount: 1}, + {name: "unknown mac does not match", macFilter: "11:22:33:44:55:66", expectedCount: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "", tt.macFilter) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } +} + +func TestSqlStore_GetAccountPeersWithExpiration(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectedCount int + expectedPeerIDs []string + }{ + { + name: "should retrieve only non-expired peers with expiration enabled", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectedCount: 1, + expectedPeerIDs: []string{"notexpired01"}, + }, + { + name: "should return no peers with expiration for a non-existing account ID", + accountID: "nonexistent", + expectedCount: 0, + }, + { + name: "should return no peers with expiration for a empty account ID", + accountID: "", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetAccountPeersWithExpiration(context.Background(), LockingStrengthNone, tt.accountID) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + for i, peer := range peers { + assert.Equal(t, tt.expectedPeerIDs[i], peer.ID) + } + }) + } +} + +func TestSqlStore_GetAccountPeersWithExpiration_ExcludesAlreadyExpired(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + peers, err := store.GetAccountPeersWithExpiration(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + + // Verify the already-expired peer (cg05lnblo1hkg2j514p0) is not returned + for _, peer := range peers { + assert.NotEqual(t, "cg05lnblo1hkg2j514p0", peer.ID, "already expired peer should not be returned") + assert.False(t, peer.Status.LoginExpired, "returned peers should not have LoginExpired set") + } +} + +func TestSqlStore_GetAccountPeersWithInactivity(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectedCount int + }{ + { + name: "should retrieve peers with inactivity for an existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectedCount: 1, + }, + { + name: "should return no peers with inactivity for a non-existing account ID", + accountID: "nonexistent", + expectedCount: 0, + }, + { + name: "should return no peers with inactivity for an empty account ID", + accountID: "", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetAccountPeersWithInactivity(context.Background(), LockingStrengthNone, tt.accountID) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } +} + +func TestSqlStore_GetAllEphemeralPeers(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/storev1.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + peers, err := store.GetAllEphemeralPeers(context.Background(), LockingStrengthNone) + require.NoError(t, err) + require.Len(t, peers, 1) + require.True(t, peers[0].Ephemeral) +} + +func TestSqlStore_GetUserPeers(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + userID string + expectedCount int + }{ + { + name: "should retrieve peers for existing account ID and user ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + userID: "f4f6d672-63fb-11ec-90d6-0242ac120003", + expectedCount: 1, + }, + { + name: "should return no peers for non-existing account ID with existing user ID", + accountID: "nonexistent", + userID: "f4f6d672-63fb-11ec-90d6-0242ac120003", + expectedCount: 0, + }, + { + name: "should return no peers for non-existing user ID with existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + userID: "nonexistent_user", + expectedCount: 0, + }, + { + name: "should retrieve peers for another valid account ID and user ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + userID: "edafee4e-63fb-11ec-90d6-0242ac120003", + expectedCount: 3, + }, + { + name: "should return no peers for existing account ID with empty user ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + userID: "", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetUserPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.userID) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } +} + +func TestSqlStore_DeletePeer(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + peerID := "csrnkiq7qv9d8aitqd50" + + err = store.DeletePeer(context.Background(), accountID, peerID) + require.NoError(t, err) + + peer, err := store.GetPeerByID(context.Background(), LockingStrengthNone, accountID, peerID) + require.Error(t, err) + require.Nil(t, peer) +} + +func BenchmarkGetAccountPeers(b *testing.B) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", b.TempDir()) + if err != nil { + b.Fatal(err) + } + b.Cleanup(cleanup) + + numberOfPeers := 1000 + numberOfGroups := 200 + numberOfPeersPerGroup := 500 + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + peers := make([]*nbpeer.Peer, 0, numberOfPeers) + for i := 0; i < numberOfPeers; i++ { + peer := &nbpeer.Peer{ + ID: fmt.Sprintf("peer-%d", i), + AccountID: accountID, + Key: fmt.Sprintf("key-%d", i), + DNSLabel: fmt.Sprintf("peer%d.example.com", i), + IP: intToIPv4(uint32(i)), + } + err = store.AddPeerToAccount(context.Background(), peer) + if err != nil { + b.Fatalf("Failed to add peer: %v", err) + } + peers = append(peers, peer) + } + + for i := 0; i < numberOfGroups; i++ { + groupID := fmt.Sprintf("group-%d", i) + group := &types.Group{ + ID: groupID, + AccountID: accountID, + } + err = store.CreateGroup(context.Background(), group) + if err != nil { + b.Fatalf("Failed to create group: %v", err) + } + for j := 0; j < numberOfPeersPerGroup; j++ { + peerIndex := (i*numberOfPeersPerGroup + j) % numberOfPeers + err = store.AddPeerToGroup(context.Background(), accountID, peers[peerIndex].ID, groupID) + if err != nil { + b.Fatalf("Failed to add peer to group: %v", err) + } + } + } + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peers[i%numberOfPeers].ID) + if err != nil { + b.Fatal(err) + } + } +} + +func intToIPv4(n uint32) netip.Addr { + var b [4]byte + binary.BigEndian.PutUint32(b[:], n) + return netip.AddrFrom4(b) +} + +func TestSqlStore_GetUserIDByPeerKey(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + userID := "test-user-123" + peerKey := "peer-key-abc" + + peer := &nbpeer.Peer{ + ID: "test-peer-1", + Key: peerKey, + AccountID: existingAccountID, + UserID: userID, + IP: netip.AddrFrom4([4]byte{10, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::a00:1"), + DNSLabel: "test-peer-1", + } + + err = store.AddPeerToAccount(context.Background(), peer) + require.NoError(t, err) + + retrievedUserID, err := store.GetUserIDByPeerKey(context.Background(), LockingStrengthNone, peerKey) + require.NoError(t, err) + assert.Equal(t, userID, retrievedUserID) +} + +func TestSqlStore_GetUserIDByPeerKey_NotFound(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + nonExistentPeerKey := "non-existent-peer-key" + + userID, err := store.GetUserIDByPeerKey(context.Background(), LockingStrengthNone, nonExistentPeerKey) + require.Error(t, err) + assert.Equal(t, "", userID) +} + +func TestSqlStore_GetUserIDByPeerKey_NoUserID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + peerKey := "peer-key-abc" + + peer := &nbpeer.Peer{ + ID: "test-peer-1", + Key: peerKey, + AccountID: existingAccountID, + UserID: "", + IP: netip.AddrFrom4([4]byte{10, 0, 0, 1}), + IPv6: netip.MustParseAddr("fd00::a00:1"), + DNSLabel: "test-peer-1", + } + + err = store.AddPeerToAccount(context.Background(), peer) + require.NoError(t, err) + + retrievedUserID, err := store.GetUserIDByPeerKey(context.Background(), LockingStrengthNone, peerKey) + require.NoError(t, err) + assert.Equal(t, "", retrievedUserID) +} + +func TestSqlStore_ApproveAccountPeers(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + accountID := "test-account" + ctx := context.Background() + + account := newAccountWithId(ctx, accountID, "testuser", "example.com") + err := store.SaveAccount(ctx, account) + require.NoError(t, err) + + peers := []*nbpeer.Peer{ + { + ID: "peer1", + AccountID: accountID, + DNSLabel: "peer1.netbird.cloud", + Key: "peer1-key", + IP: netip.MustParseAddr("100.64.0.1"), + IPv6: netip.MustParseAddr("fd00::1"), + Status: &nbpeer.PeerStatus{ + RequiresApproval: true, + LastSeen: time.Now().UTC(), + }, + }, + { + ID: "peer2", + AccountID: accountID, + DNSLabel: "peer2.netbird.cloud", + Key: "peer2-key", + IP: netip.MustParseAddr("100.64.0.2"), + IPv6: netip.MustParseAddr("fd00::2"), + Status: &nbpeer.PeerStatus{ + RequiresApproval: true, + LastSeen: time.Now().UTC(), + }, + }, + { + ID: "peer3", + AccountID: accountID, + DNSLabel: "peer3.netbird.cloud", + Key: "peer3-key", + IP: netip.MustParseAddr("100.64.0.3"), + IPv6: netip.MustParseAddr("fd00::3"), + Status: &nbpeer.PeerStatus{ + RequiresApproval: false, + LastSeen: time.Now().UTC(), + }, + }, + } + + for _, peer := range peers { + err = store.AddPeerToAccount(ctx, peer) + require.NoError(t, err) + } + + t.Run("approve all pending peers", func(t *testing.T) { + count, err := store.ApproveAccountPeers(ctx, accountID) + require.NoError(t, err) + assert.Equal(t, 2, count) + + allPeers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "", "") + require.NoError(t, err) + + for _, peer := range allPeers { + assert.False(t, peer.Status.RequiresApproval, "peer %s should not require approval", peer.ID) + } + }) + + t.Run("no peers to approve", func(t *testing.T) { + count, err := store.ApproveAccountPeers(ctx, accountID) + require.NoError(t, err) + assert.Equal(t, 0, count) + }) + + t.Run("non-existent account", func(t *testing.T) { + count, err := store.ApproveAccountPeers(ctx, "non-existent") + require.NoError(t, err) + assert.Equal(t, 0, count) + }) + }) +} diff --git a/management/server/store/sql_store_personal_access_token.go b/management/server/store/sql_store_personal_access_token.go new file mode 100644 index 000000000..d5be3e327 --- /dev/null +++ b/management/server/store/sql_store_personal_access_token.go @@ -0,0 +1,178 @@ +package store + +import ( + "context" + "database/sql" + "errors" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/management/server/util" + "github.com/netbirdio/netbird/shared/management/status" +) + +// DeleteHashedPAT2TokenIDIndex is noop in SqlStore +func (s *SqlStore) DeleteHashedPAT2TokenIDIndex(hashedToken string) error { + return nil +} + +// DeleteTokenID2UserIDIndex is noop in SqlStore +func (s *SqlStore) DeleteTokenID2UserIDIndex(tokenID string) error { + return nil +} + +func (s *SqlStore) GetTokenIDByHashedToken(ctx context.Context, hashedToken string) (string, error) { + var token types.PersonalAccessToken + result := s.db.Take(&token, "hashed_token = ?", hashedToken) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.Errorf(status.NotFound, "account not found: index lookup failed") + } + log.WithContext(ctx).Errorf("error when getting token from the store: %s", result.Error) + return "", status.NewGetAccountFromStoreError(result.Error) + } + + return token.ID, nil +} + +func (s *SqlStore) getPersonalAccessTokens(ctx context.Context, userIDs []string) ([]types.PersonalAccessToken, error) { + if len(userIDs) == 0 { + return nil, nil + } + const query = `SELECT id, user_id, name, hashed_token, expiration_date, created_by, created_at, last_used FROM personal_access_tokens WHERE user_id = ANY($1)` + rows, err := s.pgxPool().Query(ctx, query, userIDs) + if err != nil { + return nil, err + } + pats, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (types.PersonalAccessToken, error) { + var pat types.PersonalAccessToken + var expirationDate, lastUsed, createdAt sql.NullTime + err := row.Scan(&pat.ID, &pat.UserID, &pat.Name, &pat.HashedToken, &expirationDate, &pat.CreatedBy, &createdAt, &lastUsed) + if err == nil { + if expirationDate.Valid { + pat.ExpirationDate = &expirationDate.Time + } + if createdAt.Valid { + pat.CreatedAt = createdAt.Time + } + if lastUsed.Valid { + pat.LastUsed = &lastUsed.Time + } + } + return pat, err + }) + if err != nil { + return nil, err + } + return pats, nil +} + +// GetPATByHashedToken returns a PersonalAccessToken by its hashed token. +func (s *SqlStore) GetPATByHashedToken(ctx context.Context, lockStrength LockingStrength, hashedToken string) (*types.PersonalAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var pat types.PersonalAccessToken + result := tx.Take(&pat, "hashed_token = ?", hashedToken) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewPATNotFoundError(hashedToken) + } + log.WithContext(ctx).Errorf("failed to get pat by hash from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get pat by hash from store") + } + + return &pat, nil +} + +// GetPATByID retrieves a personal access token by its ID and user ID. +func (s *SqlStore) GetPATByID(ctx context.Context, lockStrength LockingStrength, userID string, patID string) (*types.PersonalAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var pat types.PersonalAccessToken + result := tx. + Take(&pat, "id = ? AND user_id = ?", patID, userID) + if err := result.Error; err != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewPATNotFoundError(patID) + } + log.WithContext(ctx).Errorf("failed to get pat from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get pat from store") + } + + return &pat, nil +} + +// GetUserPATs retrieves personal access tokens for a user. +func (s *SqlStore) GetUserPATs(ctx context.Context, lockStrength LockingStrength, userID string) ([]*types.PersonalAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var pats []*types.PersonalAccessToken + result := tx.Find(&pats, "user_id = ?", userID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get user pat's from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get user pat's from store") + } + + return pats, nil +} + +// MarkPATUsed marks a personal access token as used. +func (s *SqlStore) MarkPATUsed(ctx context.Context, patID string) error { + patCopy := types.PersonalAccessToken{ + LastUsed: util.ToPtr(time.Now().UTC()), + } + + fieldsToUpdate := []string{"last_used"} + result := s.db.Select(fieldsToUpdate). + Where(idQueryCondition, patID).Updates(&patCopy) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to mark pat as used: %s", result.Error) + return status.Errorf(status.Internal, "failed to mark pat as used") + } + + if result.RowsAffected == 0 { + return status.NewPATNotFoundError(patID) + } + + return nil +} + +// SavePAT saves a personal access token to the database. +func (s *SqlStore) SavePAT(ctx context.Context, pat *types.PersonalAccessToken) error { + result := s.db.Save(pat) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to save pat to the store: %s", err) + return status.Errorf(status.Internal, "failed to save pat to store") + } + + return nil +} + +// DeletePAT deletes a personal access token from the database. +func (s *SqlStore) DeletePAT(ctx context.Context, userID, patID string) error { + result := s.db.Delete(&types.PersonalAccessToken{}, "user_id = ? AND id = ?", userID, patID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to delete pat from the store: %s", err) + return status.Errorf(status.Internal, "failed to delete pat from store") + } + + if result.RowsAffected == 0 { + return status.NewPATNotFoundError(patID) + } + + return nil +} diff --git a/management/server/store/sql_store_personal_access_token_test.go b/management/server/store/sql_store_personal_access_token_test.go new file mode 100644 index 000000000..f40e0d9c6 --- /dev/null +++ b/management/server/store/sql_store_personal_access_token_test.go @@ -0,0 +1,186 @@ +package store + +import ( + "context" + "os" + "runtime" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/management/server/util" + "github.com/netbirdio/netbird/shared/management/status" +) + +func Test_GetTokenIDByHashedToken(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("The SQLite store is not properly supported by Windows yet") + } + + runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { + hashed := "SoMeHaShEdToKeN" + id := "9dj38s35-63fb-11ec-90d6-0242ac120003" + + token, err := store.GetTokenIDByHashedToken(context.Background(), hashed) + require.NoError(t, err) + require.Equal(t, id, token) + + _, err = store.GetTokenIDByHashedToken(context.Background(), "non-existing-hash") + require.Error(t, err) + parsedErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") + }) +} + +func TestPostgresql_GetTokenIDByHashedToken(t *testing.T) { + if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { + t.Skip("skip CI tests on darwin and windows") + } + + t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + hashed := "SoMeHaShEdToKeN" + id := "9dj38s35-63fb-11ec-90d6-0242ac120003" + + token, err := store.GetTokenIDByHashedToken(context.Background(), hashed) + require.NoError(t, err) + require.Equal(t, id, token) +} + +func TestSqlStore_GetPATByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" + + tests := []struct { + name string + patID string + expectError bool + }{ + { + name: "retrieve existing PAT", + patID: "9dj38s35-63fb-11ec-90d6-0242ac120003", + expectError: false, + }, + { + name: "retrieve non-existing PAT", + patID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty PAT ID", + patID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + pat, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, tt.patID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, pat) + } else { + require.NoError(t, err) + require.NotNil(t, pat) + require.Equal(t, tt.patID, pat.ID) + } + }) + } +} + +func TestSqlStore_GetUserPATs(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + userPATs, err := store.GetUserPATs(context.Background(), LockingStrengthNone, "f4f6d672-63fb-11ec-90d6-0242ac120003") + require.NoError(t, err) + require.Len(t, userPATs, 1) +} + +func TestSqlStore_GetPATByHashedToken(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + pat, err := store.GetPATByHashedToken(context.Background(), LockingStrengthNone, "SoMeHaShEdToKeN") + require.NoError(t, err) + require.Equal(t, "9dj38s35-63fb-11ec-90d6-0242ac120003", pat.ID) +} + +func TestSqlStore_MarkPATUsed(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" + patID := "9dj38s35-63fb-11ec-90d6-0242ac120003" + + err = store.MarkPATUsed(context.Background(), patID) + require.NoError(t, err) + + pat, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, patID) + require.NoError(t, err) + now := time.Now().UTC() + require.WithinRange(t, pat.LastUsed.UTC(), now.Add(-15*time.Second), now, "LastUsed should be within 1 second of now") +} + +func TestSqlStore_SavePAT(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + userID := "edafee4e-63fb-11ec-90d6-0242ac120003" + + pat := &types.PersonalAccessToken{ + ID: "pat-id", + UserID: userID, + Name: "token", + HashedToken: "SoMeHaShEdToKeN", + ExpirationDate: util.ToPtr(time.Now().UTC().Add(12 * time.Hour)), + CreatedBy: userID, + CreatedAt: time.Now().UTC().Add(time.Hour), + LastUsed: util.ToPtr(time.Now().UTC().Add(-15 * time.Minute)), + } + err = store.SavePAT(context.Background(), pat) + require.NoError(t, err) + + savePAT, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, pat.ID) + require.NoError(t, err) + require.Equal(t, pat.ID, savePAT.ID) + require.Equal(t, pat.UserID, savePAT.UserID) + require.Equal(t, pat.HashedToken, savePAT.HashedToken) + require.Equal(t, pat.CreatedBy, savePAT.CreatedBy) + require.WithinDurationf(t, pat.GetExpirationDate(), savePAT.ExpirationDate.UTC(), time.Millisecond, "ExpirationDate should be equal") + require.WithinDurationf(t, pat.CreatedAt, savePAT.CreatedAt.UTC(), time.Millisecond, "CreatedAt should be equal") + require.WithinDurationf(t, pat.GetLastUsed(), savePAT.LastUsed.UTC(), time.Millisecond, "LastUsed should be equal") +} + +func TestSqlStore_DeletePAT(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" + patID := "9dj38s35-63fb-11ec-90d6-0242ac120003" + + err = store.DeletePAT(context.Background(), userID, patID) + require.NoError(t, err) + + pat, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, patID) + require.Error(t, err) + require.Nil(t, pat) +} diff --git a/management/server/store/sql_store_policy.go b/management/server/store/sql_store_policy.go new file mode 100644 index 000000000..95e80e712 --- /dev/null +++ b/management/server/store/sql_store_policy.go @@ -0,0 +1,151 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getPolicies(ctx context.Context, accountID string) ([]*types.Policy, error) { + const query = `SELECT id, account_id, public_id, name, description, enabled, source_posture_checks FROM policies WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + policies, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*types.Policy, error) { + var p types.Policy + var checks []byte + var enabled sql.NullBool + err := row.Scan(&p.ID, &p.AccountID, &p.PublicID, &p.Name, &p.Description, &enabled, &checks) + if err == nil { + if enabled.Valid { + p.Enabled = enabled.Bool + } + if checks != nil { + _ = json.Unmarshal(checks, &p.SourcePostureChecks) + } + } + return &p, err + }) + if err != nil { + return nil, err + } + return policies, nil +} + +// GetAccountPolicies retrieves policies for an account. +func (s *SqlStore) GetAccountPolicies(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.Policy, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var policies []*types.Policy + result := tx. + Preload(clause.Associations).Find(&policies, accountIDCondition, accountID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get policies from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get policies from store") + } + + return policies, nil +} + +// GetPolicyByID retrieves a policy by its ID and account ID. +func (s *SqlStore) GetPolicyByID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*types.Policy, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var policy *types.Policy + + result := tx.Preload(clause.Associations). + Take(&policy, accountAndIDQueryCondition, accountID, policyID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewPolicyNotFoundError(policyID) + } + log.WithContext(ctx).Errorf("failed to get policy from store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get policy from store") + } + + return policy, nil +} + +// GetPolicyByIDOrPublicID retrieves a policy by either its ID or its PublicID. Peers report +// whichever of the two the network map they were served carries, so callers resolving a +// peer-reported reference cannot know upfront which namespace it belongs to. +func (s *SqlStore) GetPolicyByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*types.Policy, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var policy *types.Policy + + result := tx.Preload(clause.Associations). + Take(&policy, accountAndAnyIDQueryCondition, accountID, policyID, policyID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewPolicyNotFoundError(policyID) + } + log.WithContext(ctx).Errorf("failed to get policy from store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get policy from store") + } + + return policy, nil +} + +func (s *SqlStore) CreatePolicy(ctx context.Context, policy *types.Policy) error { + result := s.db.Create(policy) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to create policy in store: %s", result.Error) + return status.Errorf(status.Internal, "failed to create policy in store") + } + + return nil +} + +// SavePolicy saves a policy to the database. +func (s *SqlStore) SavePolicy(ctx context.Context, policy *types.Policy) error { + result := s.db.Session(&gorm.Session{FullSaveAssociations: true}).Omit("public_id").Save(policy) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to save policy to the store: %s", err) + return status.Errorf(status.Internal, "failed to save policy to store") + } + return nil +} + +func (s *SqlStore) DeletePolicy(ctx context.Context, accountID, policyID string) error { + return s.transaction(ctx, func(tx *gorm.DB) error { + if err := tx.Where("policy_id = ?", policyID).Delete(&types.PolicyRule{}).Error; err != nil { + return fmt.Errorf("delete policy rules: %w", err) + } + + result := tx. + Where(accountAndIDQueryCondition, accountID, policyID). + Delete(&types.Policy{}) + + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to delete policy from store: %s", err) + return status.Errorf(status.Internal, "failed to delete policy from store") + } + + if result.RowsAffected == 0 { + return status.NewPolicyNotFoundError(policyID) + } + + return nil + }) +} diff --git a/management/server/store/sql_store_policy_rule.go b/management/server/store/sql_store_policy_rule.go new file mode 100644 index 000000000..f822788ac --- /dev/null +++ b/management/server/store/sql_store_policy_rule.go @@ -0,0 +1,88 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getPolicyRules(ctx context.Context, policyIDs []string) ([]*types.PolicyRule, error) { + if len(policyIDs) == 0 { + return nil, nil + } + const query = `SELECT id, policy_id, name, description, enabled, action, destinations, destination_resource, sources, source_resource, bidirectional, protocol, ports, port_ranges, authorized_groups, authorized_user FROM policy_rules WHERE policy_id = ANY($1)` + rows, err := s.pgxPool().Query(ctx, query, policyIDs) + if err != nil { + return nil, err + } + rules, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*types.PolicyRule, error) { + var r types.PolicyRule + var dest, destRes, sources, sourceRes, ports, portRanges, authorizedGroups []byte + var enabled, bidirectional sql.NullBool + var authorizedUser sql.NullString + err := row.Scan(&r.ID, &r.PolicyID, &r.Name, &r.Description, &enabled, &r.Action, &dest, &destRes, &sources, &sourceRes, &bidirectional, &r.Protocol, &ports, &portRanges, &authorizedGroups, &authorizedUser) + if err == nil { + if enabled.Valid { + r.Enabled = enabled.Bool + } + if bidirectional.Valid { + r.Bidirectional = bidirectional.Bool + } + if dest != nil { + _ = json.Unmarshal(dest, &r.Destinations) + } + if destRes != nil { + _ = json.Unmarshal(destRes, &r.DestinationResource) + } + if sources != nil { + _ = json.Unmarshal(sources, &r.Sources) + } + if sourceRes != nil { + _ = json.Unmarshal(sourceRes, &r.SourceResource) + } + if ports != nil { + _ = json.Unmarshal(ports, &r.Ports) + } + if portRanges != nil { + _ = json.Unmarshal(portRanges, &r.PortRanges) + } + if authorizedGroups != nil { + _ = json.Unmarshal(authorizedGroups, &r.AuthorizedGroups) + } + if authorizedUser.Valid { + r.AuthorizedUser = authorizedUser.String + } + } + return &r, err + }) + if err != nil { + return nil, err + } + return rules, nil +} + +func (s *SqlStore) GetPolicyRulesByResourceID(ctx context.Context, lockStrength LockingStrength, accountID string, resourceID string) ([]*types.PolicyRule, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var policyRules []*types.PolicyRule + resourceIDPattern := `%"ID":"` + resourceID + `"%` + result := tx.Where("source_resource LIKE ? OR destination_resource LIKE ?", resourceIDPattern, resourceIDPattern). + Find(&policyRules) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get policy rules for resource id from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get policy rules for resource id from store") + } + + return policyRules, nil +} diff --git a/management/server/store/sql_store_policy_test.go b/management/server/store/sql_store_policy_test.go new file mode 100644 index 000000000..68865b184 --- /dev/null +++ b/management/server/store/sql_store_policy_test.go @@ -0,0 +1,152 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetPolicyByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + policyID string + expectError bool + }{ + { + name: "retrieve existing policy", + policyID: "cs1tnh0hhcjnqoiuebf0", + expectError: false, + }, + { + name: "retrieve non-existing policy checks", + policyID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty policy ID", + policyID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, tt.policyID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, policy) + } else { + require.NoError(t, err) + require.NotNil(t, policy) + require.Equal(t, tt.policyID, policy.ID) + } + }) + } +} + +func TestSqlStore_GetPolicyByIDOrPublicID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + policyID := "cs1tnh0hhcjnqoiuebf0" + + policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policyID) + require.NoError(t, err) + require.NotEmpty(t, policy.PublicID) + + for _, id := range []string{policyID, policy.PublicID} { + policy, err := store.GetPolicyByIDOrPublicID(context.Background(), LockingStrengthNone, accountID, id) + require.NoError(t, err) + require.Equal(t, policyID, policy.ID) + } + + policy, err = store.GetPolicyByIDOrPublicID(context.Background(), LockingStrengthNone, accountID, "non-existing") + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, policy) +} + +func TestSqlStore_CreatePolicy(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + policy := &types.Policy{ + ID: "policy-id", + AccountID: accountID, + Enabled: true, + Rules: []*types.PolicyRule{ + { + Enabled: true, + Sources: []string{"groupA"}, + Destinations: []string{"groupC"}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }, + }, + } + err = store.CreatePolicy(context.Background(), policy) + require.NoError(t, err) + + savePolicy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policy.ID) + require.NoError(t, err) + require.Equal(t, savePolicy, policy) + +} + +func TestSqlStore_SavePolicy(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + policyID := "cs1tnh0hhcjnqoiuebf0" + + policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policyID) + require.NoError(t, err) + + policy.Enabled = false + policy.Description = "policy" + policy.Rules[0].Sources = []string{"group"} + policy.Rules[0].Ports = []string{"80", "443"} + err = store.SavePolicy(context.Background(), policy) + require.NoError(t, err) + + savePolicy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policy.ID) + require.NoError(t, err) + require.Equal(t, savePolicy, policy) +} + +func TestSqlStore_DeletePolicy(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + policyID := "cs1tnh0hhcjnqoiuebf0" + + err = store.DeletePolicy(context.Background(), accountID, policyID) + require.NoError(t, err) + + policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policyID) + require.Error(t, err) + require.Nil(t, policy) +} diff --git a/management/server/store/sql_store_posture_checks.go b/management/server/store/sql_store_posture_checks.go new file mode 100644 index 000000000..d951086d4 --- /dev/null +++ b/management/server/store/sql_store_posture_checks.go @@ -0,0 +1,137 @@ +package store + +import ( + "context" + "encoding/json" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/posture" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*posture.Checks, error) { + const query = `SELECT id, account_id, public_id, name, description, checks FROM posture_checks WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + checks, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (*posture.Checks, error) { + var c posture.Checks + var checksDef []byte + err := row.Scan(&c.ID, &c.AccountID, &c.PublicID, &c.Name, &c.Description, &checksDef) + if err == nil && checksDef != nil { + _ = json.Unmarshal(checksDef, &c.Checks) + } + return &c, err + }) + if err != nil { + return nil, err + } + return checks, nil +} + +func (s *SqlStore) GetPostureCheckByChecksDefinition(accountID string, checks *posture.ChecksDefinition) (*posture.Checks, error) { + definitionJSON, err := json.Marshal(checks) + if err != nil { + return nil, err + } + + var postureCheck posture.Checks + err = s.db.Where("account_id = ? AND checks = ?", accountID, string(definitionJSON)).Take(&postureCheck).Error + if err != nil { + return nil, err + } + + return &postureCheck, nil +} + +// GetAccountPostureChecks retrieves posture checks for an account. +func (s *SqlStore) GetAccountPostureChecks(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*posture.Checks, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var postureChecks []*posture.Checks + result := tx.Find(&postureChecks, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get posture checks from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get posture checks from store") + } + + return postureChecks, nil +} + +// GetPostureChecksByID retrieves posture checks by their ID and account ID. +func (s *SqlStore) GetPostureChecksByID(ctx context.Context, lockStrength LockingStrength, accountID, postureChecksID string) (*posture.Checks, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var postureCheck *posture.Checks + result := tx. + Take(&postureCheck, accountAndIDQueryCondition, accountID, postureChecksID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewPostureChecksNotFoundError(postureChecksID) + } + log.WithContext(ctx).Errorf("failed to get posture check from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get posture check from store") + } + + return postureCheck, nil +} + +// GetPostureChecksByIDs retrieves posture checks by their IDs and account ID. +func (s *SqlStore) GetPostureChecksByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, postureChecksIDs []string) (map[string]*posture.Checks, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var postureChecks []*posture.Checks + result := tx.Find(&postureChecks, accountAndIDsQueryCondition, accountID, postureChecksIDs) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get posture checks by ID's from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get posture checks by ID's from store") + } + + postureChecksMap := make(map[string]*posture.Checks) + for _, postureCheck := range postureChecks { + postureChecksMap[postureCheck.ID] = postureCheck + } + + return postureChecksMap, nil +} + +// SavePostureChecks saves a posture checks to the database. +func (s *SqlStore) SavePostureChecks(ctx context.Context, postureCheck *posture.Checks) error { + result := s.db.Save(postureCheck) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save posture checks to store: %s", result.Error) + return status.Errorf(status.Internal, "failed to save posture checks to store") + } + + return nil +} + +// DeletePostureChecks deletes a posture checks from the database. +func (s *SqlStore) DeletePostureChecks(ctx context.Context, accountID, postureChecksID string) error { + result := s.db.Delete(&posture.Checks{}, accountAndIDQueryCondition, accountID, postureChecksID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete posture checks from store: %s", result.Error) + return status.Errorf(status.Internal, "failed to delete posture checks from store") + } + + if result.RowsAffected == 0 { + return status.NewPostureChecksNotFoundError(postureChecksID) + } + + return nil +} diff --git a/management/server/store/sql_store_posture_checks_test.go b/management/server/store/sql_store_posture_checks_test.go new file mode 100644 index 000000000..7f0511517 --- /dev/null +++ b/management/server/store/sql_store_posture_checks_test.go @@ -0,0 +1,188 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/posture" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetPostureChecksByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + postureChecksID string + expectError bool + }{ + { + name: "retrieve existing posture checks", + postureChecksID: "csplshq7qv948l48f7t0", + expectError: false, + }, + { + name: "retrieve non-existing posture checks", + postureChecksID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty posture checks ID", + postureChecksID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + postureChecks, err := store.GetPostureChecksByID(context.Background(), LockingStrengthNone, accountID, tt.postureChecksID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, postureChecks) + } else { + require.NoError(t, err) + require.NotNil(t, postureChecks) + require.Equal(t, tt.postureChecksID, postureChecks.ID) + } + }) + } +} + +func TestSqlStore_GetPostureChecksByIDs(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + postureCheckIDs []string + expectedCount int + }{ + { + name: "retrieve existing posture checks by existing IDs", + postureCheckIDs: []string{"csplshq7qv948l48f7t0", "cspnllq7qv95uq1r4k90"}, + expectedCount: 2, + }, + { + name: "empty posture check IDs list", + postureCheckIDs: []string{}, + expectedCount: 0, + }, + { + name: "non-existing posture check IDs", + postureCheckIDs: []string{"nonexistent1", "nonexistent2"}, + expectedCount: 0, + }, + { + name: "mixed existing and non-existing posture check IDs", + postureCheckIDs: []string{"cspnllq7qv95uq1r4k90", "nonexistent"}, + expectedCount: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + groups, err := store.GetPostureChecksByIDs(context.Background(), LockingStrengthNone, accountID, tt.postureCheckIDs) + require.NoError(t, err) + require.Len(t, groups, tt.expectedCount) + }) + } +} + +func TestSqlStore_SavePostureChecks(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + postureChecks := &posture.Checks{ + ID: "posture-checks-id", + AccountID: accountID, + Checks: posture.ChecksDefinition{ + NBVersionCheck: &posture.NBVersionCheck{ + MinVersion: "0.31.0", + }, + OSVersionCheck: &posture.OSVersionCheck{ + Ios: &posture.MinVersionCheck{ + MinVersion: "13.0.1", + }, + Linux: &posture.MinKernelVersionCheck{ + MinKernelVersion: "5.3.3-dev", + }, + }, + GeoLocationCheck: &posture.GeoLocationCheck{ + Locations: []posture.Location{ + { + CountryCode: "DE", + CityName: "Berlin", + }, + }, + Action: posture.CheckActionAllow, + }, + }, + } + err = store.SavePostureChecks(context.Background(), postureChecks) + require.NoError(t, err) + + savePostureChecks, err := store.GetPostureChecksByID(context.Background(), LockingStrengthNone, accountID, "posture-checks-id") + require.NoError(t, err) + require.Equal(t, savePostureChecks, postureChecks) +} + +func TestSqlStore_DeletePostureChecks(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + tests := []struct { + name string + postureChecksID string + expectError bool + }{ + { + name: "delete existing posture checks", + postureChecksID: "csplshq7qv948l48f7t0", + expectError: false, + }, + { + name: "delete non-existing posture checks", + postureChecksID: "non-existing-posture-checks-id", + expectError: true, + }, + { + name: "delete with empty posture checks ID", + postureChecksID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err = store.DeletePostureChecks(context.Background(), accountID, tt.postureChecksID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + } else { + require.NoError(t, err) + group, err := store.GetPostureChecksByID(context.Background(), LockingStrengthNone, accountID, tt.postureChecksID) + require.Error(t, err) + require.Nil(t, group) + } + }) + } +} diff --git a/management/server/store/sql_store_proxy.go b/management/server/store/sql_store_proxy.go new file mode 100644 index 000000000..bdccd282c --- /dev/null +++ b/management/server/store/sql_store_proxy.go @@ -0,0 +1,490 @@ +package store + +import ( + "context" + "errors" + "fmt" + "time" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" + "github.com/netbirdio/netbird/shared/management/status" +) + +// GetProxyMetrics aggregates per-cluster + per-proxy counts for the +// self-hosted telemetry payload. Single round-trip via conditional +// aggregations so a large proxies table doesn't fan out into multiple +// queries. +func (s *SqlStore) GetProxyMetrics(ctx context.Context) (ProxyMetrics, error) { + var m ProxyMetrics + activeCutoff := time.Now().Add(-proxyActiveThreshold) + + // COUNT(DISTINCT ... CASE WHEN ...) is portable across sqlite/postgres + // (MySQL too) and keeps the round-trip to one. proxy.StatusConnected + // is the same string the cluster-capability queries use; the active + // window matches the cluster-capability semantics (only proxies + // heartbeating within ~2 * heartbeat interval count as connected). + row := s.db.WithContext(ctx). + Model(&proxy.Proxy{}). + Select( + "COUNT(DISTINCT cluster_address) AS clusters, "+ + "COUNT(DISTINCT CASE WHEN account_id IS NOT NULL THEN cluster_address END) AS clusters_byop, "+ + "COUNT(DISTINCT CASE WHEN private = ? THEN cluster_address END) AS clusters_private, "+ + "COUNT(*) AS proxies, "+ + "COUNT(CASE WHEN status = ? AND last_seen > ? THEN 1 END) AS proxies_connected", + true, + proxy.StatusConnected, + activeCutoff, + ). + Row() + if err := row.Scan(&m.Clusters, &m.ClustersBYOP, &m.ClustersPrivate, &m.Proxies, &m.ProxiesConnected); err != nil { + return ProxyMetrics{}, fmt.Errorf("scan proxy metrics: %w", err) + } + return m, nil +} + +// SaveProxy saves or updates a proxy in the database +func (s *SqlStore) SaveProxy(ctx context.Context, p *proxy.Proxy) error { + result := s.db.Save(p) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save proxy: %v", result.Error) + return status.Errorf(status.Internal, "failed to save proxy") + } + return nil +} + +// DisconnectProxy marks a proxy as disconnected only if the session ID matches. +// This prevents a slow-to-close old session from overwriting a newer reconnection. +func (s *SqlStore) DisconnectProxy(ctx context.Context, proxyID, sessionID string) error { + now := time.Now() + result := s.db. + Model(&proxy.Proxy{}). + Where("id = ? AND session_id = ?", proxyID, sessionID). + Updates(map[string]any{ + "status": proxy.StatusDisconnected, + "disconnected_at": now, + "last_seen": now, + }) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to disconnect proxy %s session %s: %v", proxyID, sessionID, result.Error) + return status.Errorf(status.Internal, "failed to disconnect proxy") + } + if result.RowsAffected == 0 { + log.WithContext(ctx).Debugf("proxy %s session %s: no row updated (superseded by newer session)", proxyID, sessionID) + } + return nil +} + +// GetAllProxies returns all reverse proxy instance rows. +func (s *SqlStore) GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error) { + var proxies []*proxy.Proxy + result := s.db.Order("cluster_address, id").Find(&proxies) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get proxies: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get proxies") + } + return proxies, nil +} + +// DisconnectAllProxies force-marks every proxy that is not already disconnected +// as disconnected, regardless of session ID. Unlike DisconnectProxy it is not +// session-guarded: it is an administrative repair helper, not part of the +// connection lifecycle. last_seen is left untouched so the stale-proxy reaper +// keeps working off the real last heartbeat. Returns the number of proxies updated. +func (s *SqlStore) DisconnectAllProxies(ctx context.Context) (int64, error) { + result := s.db. + Model(&proxy.Proxy{}). + Where("status != ?", proxy.StatusDisconnected). + Updates(map[string]any{ + "status": proxy.StatusDisconnected, + "disconnected_at": time.Now(), + }) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to disconnect all proxies: %v", result.Error) + return 0, status.Errorf(status.Internal, "failed to disconnect all proxies") + } + return result.RowsAffected, nil +} + +// UpdateProxyHeartbeat updates the last_seen timestamp for the proxy's current session. +func (s *SqlStore) UpdateProxyHeartbeat(ctx context.Context, p *proxy.Proxy) error { + now := time.Now() + + result := s.db. + Model(&proxy.Proxy{}). + Where("id = ? AND session_id = ?", p.ID, p.SessionID). + Updates(map[string]any{ + "last_seen": now, + "status": proxy.StatusConnected, + "disconnected_at": nil, + }) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update proxy heartbeat: %v", result.Error) + return status.Errorf(status.Internal, "failed to update proxy heartbeat") + } + + if result.RowsAffected == 0 { + p.LastSeen = now + p.ConnectedAt = &now + p.Status = proxy.StatusConnected + if err := s.db.Create(p).Error; err != nil { + log.WithContext(ctx).Debugf("proxy %s session %s: heartbeat fallback insert skipped: %v", p.ID, p.SessionID, err) + } + } + + return nil +} + +// GetActiveProxyClusterAddresses returns the unique cluster addresses of active +// shared proxies (those without an account scope). BYOP cluster addresses are +// excluded; use GetActiveProxyClusterAddressesForAccount to retrieve them. +func (s *SqlStore) GetActiveProxyClusterAddresses(ctx context.Context) ([]string, error) { + var addresses []string + + result := s.db. + Model(&proxy.Proxy{}). + Where("account_id IS NULL AND status = ? AND last_seen > ?", proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). + Distinct("cluster_address"). + Pluck("cluster_address", &addresses) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get active proxy cluster addresses: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get active proxy cluster addresses") + } + + return addresses, nil +} + +func (s *SqlStore) GetActiveProxyClusterAddressesForAccount(ctx context.Context, accountID string) ([]string, error) { + var addresses []string + + result := s.db. + Model(&proxy.Proxy{}). + Where("account_id = ? AND status = ? AND last_seen > ?", accountID, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). + Distinct("cluster_address"). + Pluck("cluster_address", &addresses) + + if result.Error != nil { + return nil, status.Errorf(status.Internal, "failed to get active proxy cluster addresses for account") + } + + return addresses, nil +} + +func (s *SqlStore) GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error) { + var p proxy.Proxy + result := s.db.Where("account_id = ?", accountID).Take(&p) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "proxy not found for account") + } + return nil, status.Errorf(status.Internal, "get proxy by account ID: %v", result.Error) + } + return &p, nil +} + +func (s *SqlStore) CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error) { + var count int64 + result := s.db.Model(&proxy.Proxy{}).Where("account_id = ?", accountID).Count(&count) + if result.Error != nil { + return 0, status.Errorf(status.Internal, "count proxies by account ID: %v", result.Error) + } + return count, nil +} + +// HasActiveProxyAtClusterAddress reports whether any proxy — shared or +// account-scoped — is currently active at the given cluster address, using +// the same connected-within-threshold window as the other active-proxy +// queries. Backs the agent-network settings delete guard: settings cannot be +// deleted while a proxy declares the endpoint hostname as its address. +// +// The comparison folds case on both sides: the caller passes a normalized +// (lowercase) hostname, but proxies declare their cluster address verbatim +// and Connect stores it unchanged, so on case-sensitive collations a proxy +// declaring "GW.Example.com" would otherwise slip past the guard. Hostnames +// are case-insensitive per RFC 4343; the guard must be too. +func (s *SqlStore) HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error) { + var count int64 + result := s.db. + Model(&proxy.Proxy{}). + Where("LOWER(cluster_address) = LOWER(?) AND status = ? AND last_seen > ?", clusterAddress, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). + Count(&count) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to count active proxies at cluster address: %v", result.Error) + return false, status.Errorf(status.Internal, "failed to count active proxies at cluster address") + } + return count > 0, nil +} + +func (s *SqlStore) IsClusterAddressConflicting(ctx context.Context, clusterAddress, accountID string) (bool, error) { + var count int64 + result := s.db. + Model(&proxy.Proxy{}). + Where("cluster_address = ? AND (account_id IS NULL OR account_id != ?)", clusterAddress, accountID). + Count(&count) + if result.Error != nil { + return false, status.Errorf(status.Internal, "check cluster address conflict: %v", result.Error) + } + return count > 0, nil +} + +// HasForeignAccountProxyAtHost reports whether a proxy owned by a different +// account declares this host. Shared proxies (account_id IS NULL) are not +// foreign: a shared cluster is what most accounts pin their agent network +// gateway to. The match folds case because proxies declare their address as +// the operator spelled it while the caller's host is normalised; that costs a +// scan of the proxies table, taken once per account when its gateway is +// bootstrapped, not on the per-connect path IsClusterAddressConflicting serves. +func (s *SqlStore) HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) { + var count int64 + result := s.db. + Model(&proxy.Proxy{}). + Where("LOWER(cluster_address) = LOWER(?) AND account_id IS NOT NULL AND account_id != ?", host, accountID). + Count(&count) + if result.Error != nil { + return false, status.Errorf(status.Internal, "check proxy host ownership: %v", result.Error) + } + return count > 0, nil +} + +func (s *SqlStore) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error { + result := s.db. + Where("cluster_address = ? AND account_id = ?", clusterAddress, accountID). + Delete(&proxy.Proxy{}) + if result.Error != nil { + return status.Errorf(status.Internal, "delete account cluster: %v", result.Error) + } + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "cluster not found") + } + return nil +} + +// GetProxyClusters returns every cluster the account can see (shared +// plus its own BYOP), regardless of whether any proxy in the cluster +// is currently heartbeating. Online and ConnectedProxies are derived +// from the 2-min active window so the dashboard can render offline +// clusters distinctly; the 1-hour heartbeat reaper still removes rows +// that go quiet for too long. +// +// AccountOwned is determined by whether any proxy row in the group +// carries a non-NULL account_id; the caller maps that to Cluster.Type. +// Capability flags are NOT filled here — the handler enriches them via +// the per-cluster capability lookups. +func (s *SqlStore) GetProxyClusters(ctx context.Context, accountID string) ([]proxy.Cluster, error) { + activeCutoff := time.Now().Add(-proxyActiveThreshold) + + type clusterRow struct { + ID string + Address string + ConnectedProxies int + Online bool + AccountOwned bool + } + + var rows []clusterRow + result := s.db.Model(&proxy.Proxy{}). + Select( + "MIN(id) AS id, "+ + "cluster_address AS address, "+ + // COUNT(CASE WHEN ... THEN 1 END) counts only non-NULL — i.e. only + // rows that satisfy the predicate — so it works portably across + // sqlite/postgres/mysql without dialect-specific FILTER syntax. + "COUNT(CASE WHEN status = ? AND last_seen > ? THEN 1 END) AS connected_proxies, "+ + // MAX(CASE …) > 0 expresses BOOL_OR in a way Postgres tolerates + // (Postgres can't MAX a boolean column). + "MAX(CASE WHEN status = ? AND last_seen > ? THEN 1 ELSE 0 END) > 0 AS online, "+ + "MAX(CASE WHEN account_id IS NOT NULL THEN 1 ELSE 0 END) > 0 AS account_owned", + proxy.StatusConnected, activeCutoff, + proxy.StatusConnected, activeCutoff, + ). + Where("account_id IS NULL OR account_id = ?", accountID). + Group("cluster_address"). + Scan(&rows) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get proxy clusters: %v", result.Error) + return nil, status.Errorf(status.Internal, "get proxy clusters") + } + + clusters := make([]proxy.Cluster, 0, len(rows)) + for _, r := range rows { + c := proxy.Cluster{ + ID: r.ID, + Address: r.Address, + Online: r.Online, + ConnectedProxies: r.ConnectedProxies, + } + if r.AccountOwned { + c.Type = proxy.ClusterTypeAccount + } else { + c.Type = proxy.ClusterTypeShared + } + clusters = append(clusters, c) + } + + return clusters, nil +} + +// proxyActiveThreshold is the maximum age of a heartbeat for a proxy to be +// considered active. Must be at least 2x the heartbeat interval (1 min). +const proxyActiveThreshold = 2 * time.Minute + +var validCapabilityColumns = map[string]struct{}{ + "supports_custom_ports": {}, + "require_subdomain": {}, + "supports_crowdsec": {}, + "private": {}, +} + +// GetClusterSupportsCustomPorts returns whether any active proxy in the cluster +// supports custom ports. Returns nil when no proxy reported the capability. +func (s *SqlStore) GetClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool { + return s.getClusterCapability(ctx, clusterAddr, "supports_custom_ports") +} + +// GetClusterRequireSubdomain returns whether any active proxy in the cluster +// requires a subdomain. Returns nil when no proxy reported the capability. +func (s *SqlStore) GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { + return s.getClusterCapability(ctx, clusterAddr, "require_subdomain") +} + +// GetClusterSupportsPrivate reports whether any active proxy in the cluster +// has the private capability (nil = unreported). +func (s *SqlStore) GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool { + return s.getClusterCapability(ctx, clusterAddr, "private") +} + +// GetClusterAllProxiesPrivate reports whether every active proxy in the cluster +// has the private capability. Returns nil when no proxy reported the capability. +// Use it where any proxy in the cluster may serve the result, since a single +// non-private proxy would serve it without the private guarantees. +func (s *SqlStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + return s.getClusterUnanimousCapability(ctx, clusterAddr, "private") +} + +// GetClusterSupportsCrowdSec returns whether all active proxies in the cluster +// have CrowdSec configured. Returns nil when no proxy reported the capability. +// Unlike other capabilities that use ANY-true (for rolling upgrades), CrowdSec +// requires unanimous support: a single unconfigured proxy would let requests +// bypass reputation checks. +func (s *SqlStore) GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool { + return s.getClusterUnanimousCapability(ctx, clusterAddr, "supports_crowdsec") +} + +// GetActiveProxyVersions returns every active proxy version in a cluster. +func (s *SqlStore) GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) { + var versions []string + err := s.db.WithContext(ctx). + Model(&proxy.Proxy{}). + Where("LOWER(cluster_address) = LOWER(?) AND status = ? AND last_seen > ?", + clusterAddr, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)). + Pluck("version", &versions).Error + if err != nil { + log.WithContext(ctx).Errorf("failed to get active proxy versions for %s: %v", clusterAddr, err) + return nil, status.Errorf(status.Internal, "get active proxy versions") + } + return versions, nil +} + +// getClusterUnanimousCapability returns an aggregated boolean capability +// requiring all active proxies in the cluster to report true. +func (s *SqlStore) getClusterUnanimousCapability(ctx context.Context, clusterAddr, column string) *bool { + if _, ok := validCapabilityColumns[column]; !ok { + log.WithContext(ctx).Errorf("invalid capability column: %s", column) + return nil + } + + var result struct { + Total int64 + Reported int64 + AllTrue bool + } + + // All active proxies must have reported the capability (no NULLs) and all + // must report true. A single unreported or false proxy means the cluster + // does not unanimously support the capability. + err := s.db.WithContext(ctx). + Model(&proxy.Proxy{}). + Select("COUNT(*) AS total, "+ + "COUNT(CASE WHEN "+column+" IS NOT NULL THEN 1 END) AS reported, "+ + "COUNT(*) > 0 AND COUNT(*) = COUNT(CASE WHEN "+column+" = true THEN 1 END) AS all_true"). + Where("cluster_address = ? AND status = ? AND last_seen > ?", + clusterAddr, "connected", time.Now().Add(-proxyActiveThreshold)). + Scan(&result).Error + if err != nil { + log.WithContext(ctx).Errorf("query cluster capability %s for %s: %v", column, clusterAddr, err) + return nil + } + + if result.Total == 0 || result.Reported == 0 { + return nil + } + + // If any proxy has not reported (NULL), we can't confirm unanimous support. + if result.Reported < result.Total { + v := false + return &v + } + + return &result.AllTrue +} + +// getClusterCapability returns an aggregated boolean capability for the given +// cluster. It checks active (connected, recently seen) proxies and returns: +// - *true if any proxy in the cluster has the capability set to true, +// - *false if at least one proxy reported but none set it to true, +// - nil if no proxy reported the capability at all. +func (s *SqlStore) getClusterCapability(ctx context.Context, clusterAddr, column string) *bool { + if _, ok := validCapabilityColumns[column]; !ok { + log.WithContext(ctx).Errorf("invalid capability column: %s", column) + return nil + } + + var result struct { + HasCapability bool + AnyTrue bool + } + + err := s.db. + WithContext(ctx). + Model(&proxy.Proxy{}). + Select("COUNT(CASE WHEN "+column+" IS NOT NULL THEN 1 END) > 0 AS has_capability, "+ + "COALESCE(MAX(CASE WHEN "+column+" = true THEN 1 ELSE 0 END), 0) = 1 AS any_true"). + Where("cluster_address = ? AND status = ? AND last_seen > ?", + clusterAddr, "connected", time.Now().Add(-proxyActiveThreshold)). + Scan(&result).Error + if err != nil { + log.WithContext(ctx).Errorf("query cluster capability %s for %s: %v", column, clusterAddr, err) + return nil + } + + if !result.HasCapability { + return nil + } + + return &result.AnyTrue +} + +// CleanupStaleProxies deletes proxies that haven't sent heartbeat in the specified duration +func (s *SqlStore) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error { + cutoffTime := time.Now().Add(-inactivityDuration) + + result := s.db. + Where("last_seen < ?", cutoffTime). + Delete(&proxy.Proxy{}) + + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to cleanup stale proxies: %v", result.Error) + return status.Errorf(status.Internal, "failed to cleanup stale proxies") + } + + if result.RowsAffected > 0 { + log.WithContext(ctx).Infof("Cleaned up %d stale proxies", result.RowsAffected) + } + + return nil +} diff --git a/management/server/store/sql_store_proxy_access_token.go b/management/server/store/sql_store_proxy_access_token.go new file mode 100644 index 000000000..b111c8f64 --- /dev/null +++ b/management/server/store/sql_store_proxy_access_token.go @@ -0,0 +1,127 @@ +package store + +import ( + "context" + "errors" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// GetProxyAccessTokenByHashedToken retrieves a proxy access token by its hashed value. +func (s *SqlStore) GetProxyAccessTokenByHashedToken(ctx context.Context, lockStrength LockingStrength, hashedToken types.HashedProxyToken) (*types.ProxyAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var token types.ProxyAccessToken + result := tx.Take(&token, "hashed_token = ?", hashedToken) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "proxy access token not found") + } + return nil, status.Errorf(status.Internal, "get proxy access token: %v", result.Error) + } + + return &token, nil +} + +// GetAllProxyAccessTokens retrieves all proxy access tokens. +func (s *SqlStore) GetAllProxyAccessTokens(ctx context.Context, lockStrength LockingStrength) ([]*types.ProxyAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var tokens []*types.ProxyAccessToken + result := tx.Find(&tokens) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "get proxy access tokens: %v", result.Error) + } + + return tokens, nil +} + +// SaveProxyAccessToken saves a proxy access token to the database. +func (s *SqlStore) SaveProxyAccessToken(ctx context.Context, token *types.ProxyAccessToken) error { + if result := s.db.Create(token); result.Error != nil { + return status.Errorf(status.Internal, "save proxy access token: %v", result.Error) + } + return nil +} + +// RevokeProxyAccessToken revokes a proxy access token by its ID. +func (s *SqlStore) RevokeProxyAccessToken(ctx context.Context, tokenID string) error { + result := s.db.Model(&types.ProxyAccessToken{}).Where(idQueryCondition, tokenID).Update("revoked", true) + if result.Error != nil { + return status.Errorf(status.Internal, "revoke proxy access token: %v", result.Error) + } + + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "proxy access token not found") + } + + return nil +} + +func (s *SqlStore) GetProxyAccessTokensByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.ProxyAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var tokens []*types.ProxyAccessToken + result := tx.Where("account_id = ?", accountID).Find(&tokens) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "get proxy access tokens by account: %v", result.Error) + } + + return tokens, nil +} + +func (s *SqlStore) IsProxyAccessTokenValid(ctx context.Context, tokenID string) (bool, error) { + token, err := s.GetProxyAccessTokenByID(ctx, LockingStrengthNone, tokenID) + if err != nil { + return false, err + } + return token.IsValid(), nil +} + +func (s *SqlStore) GetProxyAccessTokenByID(ctx context.Context, lockStrength LockingStrength, tokenID string) (*types.ProxyAccessToken, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var token types.ProxyAccessToken + result := tx.Take(&token, idQueryCondition, tokenID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "proxy access token not found") + } + return nil, status.Errorf(status.Internal, "get proxy access token by ID: %v", result.Error) + } + + return &token, nil +} + +// MarkProxyAccessTokenUsed updates the last used timestamp for a proxy access token. +func (s *SqlStore) MarkProxyAccessTokenUsed(ctx context.Context, tokenID string) error { + result := s.db.Model(&types.ProxyAccessToken{}). + Where(idQueryCondition, tokenID). + Update("last_used", time.Now().UTC()) + if result.Error != nil { + return status.Errorf(status.Internal, "mark proxy access token as used: %v", result.Error) + } + + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "proxy access token not found") + } + + return nil +} diff --git a/management/server/store/sql_store_route.go b/management/server/store/sql_store_route.go new file mode 100644 index 000000000..ec9e130b1 --- /dev/null +++ b/management/server/store/sql_store_route.go @@ -0,0 +1,152 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) getRoutes(ctx context.Context, accountID string) ([]route.Route, error) { + const query = `SELECT id, account_id, public_id, network, domains, keep_route, net_id, description, peer, peer_groups, network_type, masquerade, metric, enabled, groups, access_control_groups, skip_auto_apply FROM routes WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + routes, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (route.Route, error) { + var r route.Route + var network, domains, peerGroups, groups, accessGroups []byte + var keepRoute, masquerade, enabled, skipAutoApply sql.NullBool + var metric sql.NullInt64 + err := row.Scan(&r.ID, &r.AccountID, &r.PublicID, &network, &domains, &keepRoute, &r.NetID, &r.Description, &r.Peer, &peerGroups, &r.NetworkType, &masquerade, &metric, &enabled, &groups, &accessGroups, &skipAutoApply) + if err == nil { + if keepRoute.Valid { + r.KeepRoute = keepRoute.Bool + } + if masquerade.Valid { + r.Masquerade = masquerade.Bool + } + if enabled.Valid { + r.Enabled = enabled.Bool + } + if skipAutoApply.Valid { + r.SkipAutoApply = skipAutoApply.Bool + } + if metric.Valid { + r.Metric = int(metric.Int64) + } + if network != nil { + _ = json.Unmarshal(network, &r.Network) + } + if domains != nil { + _ = json.Unmarshal(domains, &r.Domains) + } + if peerGroups != nil { + _ = json.Unmarshal(peerGroups, &r.PeerGroups) + } + if groups != nil { + _ = json.Unmarshal(groups, &r.Groups) + } + if accessGroups != nil { + _ = json.Unmarshal(accessGroups, &r.AccessControlGroups) + } + } + return r, err + }) + if err != nil { + return nil, err + } + return routes, nil +} + +// GetAccountRoutes retrieves network routes for an account. +func (s *SqlStore) GetAccountRoutes(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*route.Route, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var routes []*route.Route + result := tx.Find(&routes, accountIDCondition, accountID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get routes from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get routes from store") + } + + return routes, nil +} + +// GetRouteByID retrieves a route by its ID and account ID. +func (s *SqlStore) GetRouteByID(ctx context.Context, lockStrength LockingStrength, accountID string, routeID string) (*route.Route, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var route *route.Route + result := tx.Take(&route, accountAndIDQueryCondition, accountID, routeID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewRouteNotFoundError(routeID) + } + log.WithContext(ctx).Errorf("failed to get route from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get route from store") + } + + return route, nil +} + +// GetRouteByIDOrPublicID retrieves a route by either its ID or its PublicID. See +// GetPolicyByIDOrPublicID for why peer-reported references need both. +func (s *SqlStore) GetRouteByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID string, routeID string) (*route.Route, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var route *route.Route + result := tx.Take(&route, accountAndAnyIDQueryCondition, accountID, routeID, routeID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewRouteNotFoundError(routeID) + } + log.WithContext(ctx).Errorf("failed to get route from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get route from store") + } + + return route, nil +} + +// SaveRoute saves a route to the database. +func (s *SqlStore) SaveRoute(ctx context.Context, route *route.Route) error { + result := s.db.Save(route) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to save route to the store: %s", err) + return status.Errorf(status.Internal, "failed to save route to store") + } + + return nil +} + +// DeleteRoute deletes a route from the database. +func (s *SqlStore) DeleteRoute(ctx context.Context, accountID, routeID string) error { + result := s.db.Delete(&route.Route{}, accountAndIDQueryCondition, accountID, routeID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to delete route from the store: %s", err) + return status.Errorf(status.Internal, "failed to delete route from store") + } + + if result.RowsAffected == 0 { + return status.NewRouteNotFoundError(routeID) + } + + return nil +} diff --git a/management/server/store/sql_store_route_test.go b/management/server/store/sql_store_route_test.go new file mode 100644 index 000000000..53e4f130e --- /dev/null +++ b/management/server/store/sql_store_route_test.go @@ -0,0 +1,165 @@ +package store + +import ( + "context" + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + + nbroute "github.com/netbirdio/netbird/route" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_GetAccountRoutes(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + expectedCount int + }{ + { + name: "retrieve routes by existing account ID", + accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", + expectedCount: 1, + }, + { + name: "non-existing account ID", + accountID: "nonexistent", + expectedCount: 0, + }, + { + name: "empty account ID", + accountID: "", + expectedCount: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + routes, err := store.GetAccountRoutes(context.Background(), LockingStrengthNone, tt.accountID) + require.NoError(t, err) + require.Len(t, routes, tt.expectedCount) + }) + } +} + +func TestSqlStore_GetRouteByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + tests := []struct { + name string + routeID string + expectError bool + }{ + { + name: "retrieve existing route", + routeID: "ct03t427qv97vmtmglog", + expectError: false, + }, + { + name: "retrieve non-existing route", + routeID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty route ID", + routeID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + route, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, tt.routeID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, route) + } else { + require.NoError(t, err) + require.NotNil(t, route) + require.Equal(t, tt.routeID, string(route.ID)) + } + }) + } +} + +func TestSqlStore_GetRouteByIDOrPublicID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + routeID := "ct03t427qv97vmtmglog" + + route, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, routeID) + require.NoError(t, err) + require.NotEmpty(t, route.PublicID) + + for _, id := range []string{routeID, route.PublicID} { + route, err := store.GetRouteByIDOrPublicID(context.Background(), LockingStrengthNone, accountID, id) + require.NoError(t, err) + require.Equal(t, routeID, string(route.ID)) + } + + route, err = store.GetRouteByIDOrPublicID(context.Background(), LockingStrengthNone, accountID, "non-existing") + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, route) +} + +func TestSqlStore_SaveRoute(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + route := &nbroute.Route{ + ID: "route-id", + AccountID: accountID, + Network: netip.MustParsePrefix("10.10.0.0/16"), + NetID: "netID", + PeerGroups: []string{"routeA"}, + NetworkType: nbroute.IPv4Network, + Masquerade: true, + Metric: 9999, + Enabled: true, + Groups: []string{"groupA"}, + AccessControlGroups: []string{}, + } + err = store.SaveRoute(context.Background(), route) + require.NoError(t, err) + + saveRoute, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, string(route.ID)) + require.NoError(t, err) + require.Equal(t, route, saveRoute) + +} + +func TestSqlStore_DeleteRoute(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + routeID := "ct03t427qv97vmtmglog" + + err = store.DeleteRoute(context.Background(), accountID, routeID) + require.NoError(t, err) + + route, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, routeID) + require.Error(t, err) + require.Nil(t, route) +} diff --git a/management/server/store/sql_store_service.go b/management/server/store/sql_store_service.go new file mode 100644 index 000000000..4f4f3546d --- /dev/null +++ b/management/server/store/sql_store_service.go @@ -0,0 +1,452 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "math" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/shared/management/status" +) + +// serviceSelectColumns and targetSelectColumns are the column lists the Postgres +// pgx read path scans. They must stay in sync with the rpservice.Service and +// rpservice.Target gorm models; TestPgxServiceColumnsMatchGorm enforces this. +const serviceSelectColumns = `id, account_id, name, domain, enabled, auth, restrictions, + meta_created_at, meta_certificate_issued_at, meta_last_renewed_at, meta_status, proxy_cluster, + pass_host_header, rewrite_redirects, session_private_key, session_public_key, + mode, listen_port, port_auto_assigned, source, source_peer, terminated, + private, access_groups` + +func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { + const serviceQuery = `SELECT ` + serviceSelectColumns + ` FROM services WHERE account_id = $1` + + serviceRows, err := s.pgxPool().Query(ctx, serviceQuery, accountID) + if err != nil { + return nil, err + } + + services, err := pgx.CollectRows(serviceRows, scanService) + if err != nil { + return nil, err + } + + if len(services) == 0 { + return services, nil + } + + serviceIDs := make([]string, len(services)) + serviceMap := make(map[string]*rpservice.Service) + for i, svc := range services { + serviceIDs[i] = svc.ID + serviceMap[svc.ID] = svc + } + + targets, err := s.getServiceTargets(ctx, serviceIDs) + if err != nil { + return nil, err + } + + for _, target := range targets { + if service, ok := serviceMap[target.ServiceID]; ok { + service.Targets = append(service.Targets, target) + } + } + + return services, nil +} + +func scanService(row pgx.CollectableRow) (*rpservice.Service, error) { + var s rpservice.Service + var auth []byte + var restrictions []byte + var accessGroups []byte + var createdAt, certIssuedAt, lastRenewedAt sql.NullTime + var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString + var mode, source, sourcePeer sql.NullString + var terminated, portAutoAssigned, private sql.NullBool + var listenPort sql.NullInt64 + err := row.Scan( + &s.ID, + &s.AccountID, + &s.Name, + &s.Domain, + &s.Enabled, + &auth, + &restrictions, + &createdAt, + &certIssuedAt, + &lastRenewedAt, + &status, + &proxyCluster, + &s.PassHostHeader, + &s.RewriteRedirects, + &sessionPrivateKey, + &sessionPublicKey, + &mode, + &listenPort, + &portAutoAssigned, + &source, + &sourcePeer, + &terminated, + &private, + &accessGroups, + ) + if err != nil { + return nil, err + } + + if auth != nil { + if err := json.Unmarshal(auth, &s.Auth); err != nil { + return nil, err + } + } + + if len(restrictions) > 0 { + if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil { + return nil, fmt.Errorf("unmarshal restrictions: %w", err) + } + } + + if len(accessGroups) > 0 { + if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { + return nil, fmt.Errorf("unmarshal access_groups: %w", err) + } + } + + if private.Valid { + s.Private = private.Bool + } + + s.Meta = serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt, status) + if proxyCluster.Valid { + s.ProxyCluster = proxyCluster.String + } + if sessionPrivateKey.Valid { + s.SessionPrivateKey = sessionPrivateKey.String + } + if sessionPublicKey.Valid { + s.SessionPublicKey = sessionPublicKey.String + } + if mode.Valid { + s.Mode = mode.String + } + if source.Valid { + s.Source = source.String + } + if sourcePeer.Valid { + s.SourcePeer = sourcePeer.String + } + if terminated.Valid { + s.Terminated = terminated.Bool + } + if portAutoAssigned.Valid { + s.PortAutoAssigned = portAutoAssigned.Bool + } + if listenPort.Valid { + if listenPort.Int64 < 0 || listenPort.Int64 > math.MaxUint16 { + return nil, fmt.Errorf("listen_port %d out of range", listenPort.Int64) + } + s.ListenPort = uint16(listenPort.Int64) + } + s.Targets = []*rpservice.Target{} + return &s, nil +} + +func serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt sql.NullTime, status sql.NullString) rpservice.Meta { + meta := rpservice.Meta{} + if createdAt.Valid { + meta.CreatedAt = createdAt.Time + } + if certIssuedAt.Valid { + t := certIssuedAt.Time + meta.CertificateIssuedAt = &t + } + if lastRenewedAt.Valid { + t := lastRenewedAt.Time + meta.LastRenewedAt = &t + } + if status.Valid { + meta.Status = status.String + } + return meta +} + +func (s *SqlStore) CreateService(ctx context.Context, service *rpservice.Service) error { + serviceCopy := service.Copy() + if err := serviceCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { + return fmt.Errorf("encrypt service data: %w", err) + } + result := s.db.Create(serviceCopy) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to create service to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to create service to store") + } + + return nil +} + +func (s *SqlStore) UpdateService(ctx context.Context, service *rpservice.Service) error { + serviceCopy := service.Copy() + if err := serviceCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { + return fmt.Errorf("encrypt service data: %w", err) + } + + // Create target type instance outside transaction to avoid variable shadowing + targetType := &rpservice.Target{} + + // Use a transaction to ensure atomic updates of the service and its targets + err := s.transaction(ctx, func(tx *gorm.DB) error { + // Delete existing targets + if err := tx.Where("service_id = ?", serviceCopy.ID).Delete(targetType).Error; err != nil { + return err + } + + // Update the service and create new targets + if err := tx.Session(&gorm.Session{FullSaveAssociations: true}).Save(serviceCopy).Error; err != nil { + return err + } + + return nil + }) + if err != nil { + log.WithContext(ctx).Errorf("failed to update service to store: %v", err) + return status.Errorf(status.Internal, "failed to update service to store") + } + + return nil +} + +func (s *SqlStore) DeleteService(ctx context.Context, accountID, serviceID string) error { + result := s.db.Delete(&rpservice.Service{}, accountAndIDQueryCondition, accountID, serviceID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete service from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete service from store") + } + + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "service %s not found", serviceID) + } + + return nil +} + +func (s *SqlStore) GetServiceByID(ctx context.Context, lockStrength LockingStrength, accountID, serviceID string) (*rpservice.Service, error) { + tx := s.db.Preload("Targets") + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var service *rpservice.Service + result := tx.Take(&service, accountAndIDQueryCondition, accountID, serviceID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "service %s not found", serviceID) + } + + log.WithContext(ctx).Errorf("failed to get service from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get service from store") + } + + if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt service data: %w", err) + } + + return service, nil +} + +func (s *SqlStore) GetServiceByDomain(ctx context.Context, domain string) (*rpservice.Service, error) { + var service *rpservice.Service + result := s.db.Preload("Targets").Where("domain = ?", domain).First(&service) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "service with domain %s not found", domain) + } + + log.WithContext(ctx).Errorf("failed to get service by domain from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get service by domain from store") + } + + if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt service data: %w", err) + } + + return service, nil +} + +func (s *SqlStore) GetServices(ctx context.Context, lockStrength LockingStrength) ([]*rpservice.Service, error) { + tx := s.db.Preload("Targets") + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var serviceList []*rpservice.Service + result := tx.Find(&serviceList) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get services from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get services from store") + } + + for _, service := range serviceList { + if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt service data: %w", err) + } + } + + return serviceList, nil +} + +func (s *SqlStore) GetAccountServices(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*rpservice.Service, error) { + tx := s.db.Preload("Targets") + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var serviceList []*rpservice.Service + result := tx.Find(&serviceList, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get services from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get services from store") + } + + for _, service := range serviceList { + if err := service.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt service data: %w", err) + } + } + + return serviceList, nil +} + +// RenewEphemeralService updates the last_renewed_at timestamp for an ephemeral service. +func (s *SqlStore) RenewEphemeralService(ctx context.Context, accountID, peerID, serviceID string) error { + result := s.db.Model(&rpservice.Service{}). + Where("id = ? AND account_id = ? AND source_peer = ? AND source = ?", serviceID, accountID, peerID, rpservice.SourceEphemeral). + Update("meta_last_renewed_at", time.Now()) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to renew ephemeral service: %v", result.Error) + return status.Errorf(status.Internal, "renew ephemeral service") + } + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "no active expose session for service %s", serviceID) + } + return nil +} + +// GetExpiredEphemeralServices returns ephemeral services whose last renewal exceeds the given TTL. +// Only the fields needed for reaping are selected. The limit parameter caps the batch size to +// avoid loading too many rows in a single tick. Rows with empty source_peer are excluded to +// skip malformed legacy data. +func (s *SqlStore) GetExpiredEphemeralServices(ctx context.Context, ttl time.Duration, limit int) ([]*rpservice.Service, error) { + cutoff := time.Now().Add(-ttl) + var services []*rpservice.Service + result := s.db. + Select("id", "account_id", "source_peer", "domain"). + Where("source = ? AND source_peer <> '' AND meta_last_renewed_at < ?", rpservice.SourceEphemeral, cutoff). + Limit(limit). + Find(&services) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get expired ephemeral services: %v", result.Error) + return nil, status.Errorf(status.Internal, "get expired ephemeral services") + } + return services, nil +} + +// CountEphemeralServicesByPeer returns the count of ephemeral services for a specific peer. +// Use LockingStrengthUpdate inside a transaction to serialize concurrent create operations. +// The locking is applied via a row-level SELECT ... FOR UPDATE (not on the aggregate) to +// stay compatible with Postgres, which disallows FOR UPDATE on COUNT(*). +func (s *SqlStore) CountEphemeralServicesByPeer(ctx context.Context, lockStrength LockingStrength, accountID, peerID string) (int64, error) { + if lockStrength == LockingStrengthNone { + var count int64 + result := s.db.Model(&rpservice.Service{}). + Where("account_id = ? AND source_peer = ? AND source = ?", accountID, peerID, rpservice.SourceEphemeral). + Count(&count) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to count ephemeral services: %v", result.Error) + return 0, status.Errorf(status.Internal, "count ephemeral services") + } + return count, nil + } + + var ids []string + result := s.db.Model(&rpservice.Service{}). + Clauses(clause.Locking{Strength: string(lockStrength)}). + Select("id"). + Where("account_id = ? AND source_peer = ? AND source = ?", accountID, peerID, rpservice.SourceEphemeral). + Pluck("id", &ids) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to count ephemeral services: %v", result.Error) + return 0, status.Errorf(status.Internal, "count ephemeral services") + } + return int64(len(ids)), nil +} + +// EphemeralServiceExists checks if an ephemeral service exists for the given peer and domain. +// Use LockingStrengthUpdate inside a transaction to serialize concurrent create operations. +func (s *SqlStore) EphemeralServiceExists(ctx context.Context, lockStrength LockingStrength, accountID, peerID, domain string) (bool, error) { + if lockStrength == LockingStrengthNone { + var count int64 + result := s.db.Model(&rpservice.Service{}). + Where("account_id = ? AND source_peer = ? AND domain = ? AND source = ?", accountID, peerID, domain, rpservice.SourceEphemeral). + Count(&count) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to check ephemeral service existence: %v", result.Error) + return false, status.Errorf(status.Internal, "check ephemeral service existence") + } + return count > 0, nil + } + + var id string + result := s.db.Model(&rpservice.Service{}). + Clauses(clause.Locking{Strength: string(lockStrength)}). + Select("id"). + Where("account_id = ? AND source_peer = ? AND domain = ? AND source = ?", accountID, peerID, domain, rpservice.SourceEphemeral). + Limit(1). + Pluck("id", &id) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to check ephemeral service existence: %v", result.Error) + return false, status.Errorf(status.Internal, "check ephemeral service existence") + } + return id != "", nil +} + +// GetServicesByClusterAndPort returns services matching the given proxy cluster, mode, and listen port. +func (s *SqlStore) GetServicesByClusterAndPort(ctx context.Context, lockStrength LockingStrength, proxyCluster string, mode string, listenPort uint16) ([]*rpservice.Service, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var services []*rpservice.Service + result := tx.Where("proxy_cluster = ? AND mode = ? AND listen_port = ?", proxyCluster, mode, listenPort).Find(&services) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "query services by cluster and port") + } + + return services, nil +} + +// GetServicesByCluster returns all services for the given proxy cluster. +func (s *SqlStore) GetServicesByCluster(ctx context.Context, lockStrength LockingStrength, proxyCluster string) ([]*rpservice.Service, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var services []*rpservice.Service + result := tx.Where("proxy_cluster = ?", proxyCluster).Find(&services) + if result.Error != nil { + return nil, status.Errorf(status.Internal, "query services by cluster") + } + return services, nil +} diff --git a/management/server/store/sql_store_service_target.go b/management/server/store/sql_store_service_target.go new file mode 100644 index 000000000..5c2b99510 --- /dev/null +++ b/management/server/store/sql_store_service_target.go @@ -0,0 +1,163 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/shared/management/status" +) + +const targetSelectColumns = `id, account_id, service_id, path, host, port, protocol, + target_id, target_type, enabled, proxy_protocol, + skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers, + direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes, + capture_content_types, agent_network, disable_access_log` + +func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) { + const targetsQuery = `SELECT ` + targetSelectColumns + ` FROM targets WHERE service_id = ANY($1)` + + rows, err := s.pgxPool().Query(ctx, targetsQuery, serviceIDs) + if err != nil { + return nil, err + } + + return pgx.CollectRows(rows, scanTarget) +} + +func scanTarget(row pgx.CollectableRow) (*rpservice.Target, error) { + var t rpservice.Target + var path sql.NullString + var pathRewrite sql.NullString + var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool + var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64 + var customHeaders, middlewares, captureContentTypes []byte + err := row.Scan( + &t.ID, + &t.AccountID, + &t.ServiceID, + &path, + &t.Host, + &t.Port, + &t.Protocol, + &t.TargetId, + &t.TargetType, + &t.Enabled, + &proxyProtocol, + &skipTLSVerify, + &requestTimeout, + &sessionIdleTimeout, + &pathRewrite, + &customHeaders, + &directUpstream, + &middlewares, + &captureMaxRequestBytes, + &captureMaxResponseBytes, + &captureContentTypes, + &agentNetwork, + &disableAccessLog, + ) + if err != nil { + return nil, err + } + if path.Valid { + t.Path = &path.String + } + + t.ProxyProtocol = proxyProtocol.Bool + t.Options.SkipTLSVerify = skipTLSVerify.Bool + t.Options.RequestTimeout = time.Duration(requestTimeout.Int64) + t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64) + t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String) + t.Options.DirectUpstream = directUpstream.Bool + t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64 + t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64 + t.Options.AgentNetwork = agentNetwork.Bool + t.Options.DisableAccessLog = disableAccessLog.Bool + + if len(customHeaders) > 0 { + if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil { + return nil, fmt.Errorf("unmarshal custom_headers: %w", err) + } + } + if len(middlewares) > 0 { + if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil { + return nil, fmt.Errorf("unmarshal middlewares: %w", err) + } + } + if len(captureContentTypes) > 0 { + if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil { + return nil, fmt.Errorf("unmarshal capture_content_types: %w", err) + } + } + return &t, nil +} + +func (s *SqlStore) DeleteTarget(ctx context.Context, accountID string, serviceID string, targetID uint) error { + result := s.db.Delete(&rpservice.Target{}, "account_id = ? AND service_id = ? AND id = ?", accountID, serviceID, targetID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete target from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete target from store") + } + + if result.RowsAffected == 0 { + return status.Errorf(status.NotFound, "target not found for service %s", serviceID) + } + + return nil +} + +func (s *SqlStore) DeleteServiceTargets(ctx context.Context, accountID string, serviceID string) error { + result := s.db.Delete(&rpservice.Target{}, "account_id = ? AND service_id = ?", accountID, serviceID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete targets from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete targets from store") + } + + return nil +} + +// GetTargetsByServiceID retrieves all targets for a given service +func (s *SqlStore) GetTargetsByServiceID(ctx context.Context, lockStrength LockingStrength, accountID string, serviceID string) ([]*rpservice.Target, error) { + var targets []*rpservice.Target + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + result := tx.Where("account_id = ? AND service_id = ?", accountID, serviceID).Find(&targets) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get targets from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get targets from store") + } + + return targets, nil +} + +func (s *SqlStore) GetServiceTargetByTargetID(ctx context.Context, lockStrength LockingStrength, accountID string, targetID string) (*rpservice.Target, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var target *rpservice.Target + result := tx.Take(&target, "account_id = ? AND target_id = ?", accountID, targetID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "service target with ID %s not found", targetID) + } + + log.WithContext(ctx).Errorf("failed to get service target from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get service target from store") + } + + return target, nil +} diff --git a/management/server/store/sql_store_setup_key.go b/management/server/store/sql_store_setup_key.go new file mode 100644 index 000000000..3857f6e8c --- /dev/null +++ b/management/server/store/sql_store_setup_key.go @@ -0,0 +1,219 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) GetAccountBySetupKey(ctx context.Context, setupKey string) (*types.Account, error) { + var key types.SetupKey + result := s.db.Select("account_id").Take(&key, GetKeyQueryCondition(s), setupKey) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewSetupKeyNotFoundError(setupKey) + } + log.WithContext(ctx).Errorf("failed to get account by setup key from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get account by setup key from store") + } + + if key.AccountID == "" { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + + return s.GetAccount(ctx, key.AccountID) +} + +func (s *SqlStore) getSetupKeys(ctx context.Context, accountID string) ([]types.SetupKey, error) { + const query = `SELECT id, account_id, key, key_secret, name, type, created_at, expires_at, updated_at, + revoked, used_times, last_used, auto_groups, usage_limit, ephemeral, allow_extra_dns_labels FROM setup_keys WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + + keys, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (types.SetupKey, error) { + var sk types.SetupKey + var autoGroups []byte + var skCreatedAt, expiresAt, updatedAt, lastUsed sql.NullTime + var revoked, ephemeral, allowExtraDNSLabels sql.NullBool + var usedTimes, usageLimit sql.NullInt64 + + err := row.Scan(&sk.Id, &sk.AccountID, &sk.Key, &sk.KeySecret, &sk.Name, &sk.Type, &skCreatedAt, + &expiresAt, &updatedAt, &revoked, &usedTimes, &lastUsed, &autoGroups, &usageLimit, &ephemeral, &allowExtraDNSLabels) + + if err == nil { + if expiresAt.Valid { + sk.ExpiresAt = &expiresAt.Time + } + if skCreatedAt.Valid { + sk.CreatedAt = skCreatedAt.Time + } + if updatedAt.Valid { + sk.UpdatedAt = updatedAt.Time + if sk.UpdatedAt.IsZero() { + sk.UpdatedAt = sk.CreatedAt + } + } + if lastUsed.Valid { + sk.LastUsed = &lastUsed.Time + } + if revoked.Valid { + sk.Revoked = revoked.Bool + } + if usedTimes.Valid { + sk.UsedTimes = int(usedTimes.Int64) + } + if usageLimit.Valid { + sk.UsageLimit = int(usageLimit.Int64) + } + if ephemeral.Valid { + sk.Ephemeral = ephemeral.Bool + } + if allowExtraDNSLabels.Valid { + sk.AllowExtraDNSLabels = allowExtraDNSLabels.Bool + } + if autoGroups != nil { + _ = json.Unmarshal(autoGroups, &sk.AutoGroups) + } else { + sk.AutoGroups = []string{} + } + } + return sk, err + }) + if err != nil { + return nil, err + } + return keys, nil +} + +func (s *SqlStore) GetAccountIDBySetupKey(ctx context.Context, setupKey string) (string, error) { + var accountID string + result := s.db.Model(&types.SetupKey{}).Select("account_id").Where(GetKeyQueryCondition(s), setupKey).Take(&accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.NewSetupKeyNotFoundError(setupKey) + } + log.WithContext(ctx).Errorf("failed to get account ID by setup key from store: %v", result.Error) + return "", status.Errorf(status.Internal, "failed to get account ID by setup key from store") + } + + if accountID == "" { + return "", status.Errorf(status.NotFound, "account not found: index lookup failed") + } + + return accountID, nil +} + +func (s *SqlStore) GetSetupKeyBySecret(ctx context.Context, lockStrength LockingStrength, key string) (*types.SetupKey, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var setupKey types.SetupKey + result := tx. + Take(&setupKey, GetKeyQueryCondition(s), key) + + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.PreconditionFailed, "setup key not found") + } + log.WithContext(ctx).Errorf("failed to get setup key by secret from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get setup key by secret from store") + } + return &setupKey, nil +} + +func (s *SqlStore) IncrementSetupKeyUsage(ctx context.Context, setupKeyID string) error { + result := s.db.Model(&types.SetupKey{}). + Where(idQueryCondition, setupKeyID). + Updates(map[string]interface{}{ + "used_times": gorm.Expr("used_times + 1"), + "last_used": time.Now(), + }) + + if result.Error != nil { + return status.Errorf(status.Internal, "issue incrementing setup key usage count: %s", result.Error) + } + + if result.RowsAffected == 0 { + return status.NewSetupKeyNotFoundError(setupKeyID) + } + + return nil +} + +// GetAccountSetupKeys retrieves setup keys for an account. +func (s *SqlStore) GetAccountSetupKeys(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.SetupKey, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var setupKeys []*types.SetupKey + result := tx. + Find(&setupKeys, accountIDCondition, accountID) + if err := result.Error; err != nil { + log.WithContext(ctx).Errorf("failed to get setup keys from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get setup keys from store") + } + + return setupKeys, nil +} + +// GetSetupKeyByID retrieves a setup key by its ID and account ID. +func (s *SqlStore) GetSetupKeyByID(ctx context.Context, lockStrength LockingStrength, accountID, setupKeyID string) (*types.SetupKey, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var setupKey *types.SetupKey + result := tx.Take(&setupKey, accountAndIDQueryCondition, accountID, setupKeyID) + if err := result.Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, status.NewSetupKeyNotFoundError(setupKeyID) + } + log.WithContext(ctx).Errorf("failed to get setup key from the store: %s", err) + return nil, status.Errorf(status.Internal, "failed to get setup key from store") + } + + return setupKey, nil +} + +// SaveSetupKey saves a setup key to the database. +func (s *SqlStore) SaveSetupKey(ctx context.Context, setupKey *types.SetupKey) error { + result := s.db.Save(setupKey) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save setup key to store: %s", result.Error) + return status.Errorf(status.Internal, "failed to save setup key to store") + } + + return nil +} + +// DeleteSetupKey deletes a setup key from the database. +func (s *SqlStore) DeleteSetupKey(ctx context.Context, accountID, keyID string) error { + result := s.db.Delete(&types.SetupKey{}, accountAndIDQueryCondition, accountID, keyID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete setup key from store: %s", result.Error) + return status.Errorf(status.Internal, "failed to delete setup key from store") + } + + if result.RowsAffected == 0 { + return status.NewSetupKeyNotFoundError(keyID) + } + + return nil +} diff --git a/management/server/store/sql_store_setup_key_test.go b/management/server/store/sql_store_setup_key_test.go new file mode 100644 index 000000000..8835ea1f9 --- /dev/null +++ b/management/server/store/sql_store_setup_key_test.go @@ -0,0 +1,103 @@ +package store + +import ( + "context" + "crypto/sha256" + b64 "encoding/base64" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/types" +) + +func TestSqlite_GetSetupKeyBySecret(t *testing.T) { + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + if err != nil { + t.Fatal(err) + } + + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + plainKey := "A2C8E62B-38F5-4553-B31E-DD66C696CEBB" + hashedKey := sha256.Sum256([]byte(plainKey)) + encodedHashedKey := b64.StdEncoding.EncodeToString(hashedKey[:]) + + _, err = store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + setupKey, err := store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) + require.NoError(t, err) + assert.Equal(t, encodedHashedKey, setupKey.Key) + assert.Equal(t, types.HiddenKey(plainKey, 4), setupKey.KeySecret) + assert.Equal(t, "bf1c8084-ba50-4ce7-9439-34653001fc3b", setupKey.AccountID) + assert.Equal(t, "Default key", setupKey.Name) +} + +func TestSqlite_incrementSetupKeyUsage(t *testing.T) { + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + if err != nil { + t.Fatal(err) + } + + existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + plainKey := "A2C8E62B-38F5-4553-B31E-DD66C696CEBB" + hashedKey := sha256.Sum256([]byte(plainKey)) + encodedHashedKey := b64.StdEncoding.EncodeToString(hashedKey[:]) + + _, err = store.GetAccount(context.Background(), existingAccountID) + require.NoError(t, err) + + setupKey, err := store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) + require.NoError(t, err) + assert.Equal(t, 0, setupKey.UsedTimes) + + err = store.IncrementSetupKeyUsage(context.Background(), setupKey.Id) + require.NoError(t, err) + + setupKey, err = store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) + require.NoError(t, err) + assert.Equal(t, 1, setupKey.UsedTimes) + + err = store.IncrementSetupKeyUsage(context.Background(), setupKey.Id) + require.NoError(t, err) + + setupKey, err = store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) + require.NoError(t, err) + assert.Equal(t, 2, setupKey.UsedTimes) +} + +func Test_DeleteSetupKeySuccessfully(t *testing.T) { + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + setupKeyID := "A2C8E62B-38F5-4553-B31E-DD66C696CEBB" + + err = store.DeleteSetupKey(context.Background(), accountID, setupKeyID) + require.NoError(t, err) + + _, err = store.GetSetupKeyByID(context.Background(), LockingStrengthNone, setupKeyID, accountID) + require.Error(t, err) +} + +func Test_DeleteSetupKeyFailsForNonExistingKey(t *testing.T) { + t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + nonExistingKeyID := "non-existing-key-id" + + err = store.DeleteSetupKey(context.Background(), accountID, nonExistingKeyID) + require.Error(t, err) +} diff --git a/management/server/store/sql_store_test.go b/management/server/store/sql_store_test.go index fbcff5257..f695a150d 100644 --- a/management/server/store/sql_store_test.go +++ b/management/server/store/sql_store_test.go @@ -2,16 +2,11 @@ package store import ( "context" - "crypto/sha256" - b64 "encoding/base64" - "encoding/binary" "fmt" "net" "net/netip" "os" - "reflect" "runtime" - "sort" "sync" "testing" "time" @@ -20,23 +15,13 @@ import ( log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" + "gorm.io/gorm/clause" nbdns "github.com/netbirdio/netbird/dns" - proxydomain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" - rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" - "github.com/netbirdio/netbird/management/internals/modules/zones" - "github.com/netbirdio/netbird/management/internals/modules/zones/records" - resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" - routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" - networkTypes "github.com/netbirdio/netbird/management/server/networks/types" nbpeer "github.com/netbirdio/netbird/management/server/peer" - "github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/types" - "github.com/netbirdio/netbird/management/server/util" nbroute "github.com/netbirdio/netbird/route" - "github.com/netbirdio/netbird/shared/management/status" - "github.com/netbirdio/netbird/shared/testing_helpers" - "github.com/netbirdio/netbird/util/crypt" ) func runTestForAllEngines(t *testing.T, testDataFile string, f func(t *testing.T, store Store)) { @@ -72,650 +57,6 @@ func Test_NewStore(t *testing.T) { }) } -func Test_SaveAccount_Large(t *testing.T) { - if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { - t.Skip("skip CI tests on darwin and windows") - } - - runTestForAllEngines(t, "", func(t *testing.T, store Store) { - runLargeTest(t, store) - }) -} - -func runLargeTest(t *testing.T, store Store) { - t.Helper() - - account := newAccountWithId(context.Background(), "account_id", "testuser", "") - groupALL, err := account.GetGroupAll() - if err != nil { - t.Fatal(err) - } - setupKey, _ := types.GenerateDefaultSetupKey() - account.SetupKeys[setupKey.Key] = setupKey - const numPerAccount = 6000 - for n := 0; n < numPerAccount; n++ { - netIP := sequentialIPv4(n) - peerID := fmt.Sprintf("%s-peer-%d", account.Id, n) - addr, _ := netip.AddrFromSlice(netIP) - - peer := &nbpeer.Peer{ - ID: peerID, - Key: peerID, - IP: addr.Unmap(), - Name: peerID, - DNSLabel: peerID, - UserID: "testuser", - Status: &nbpeer.PeerStatus{Connected: false, LastSeen: time.Now()}, - SSHEnabled: false, - } - account.Peers[peerID] = peer - group, _ := account.GetGroupAll() - group.Peers = append(group.Peers, peerID) - user := &types.User{ - Id: fmt.Sprintf("%s-user-%d", account.Id, n), - AccountID: account.Id, - } - account.Users[user.Id] = user - route := &nbroute.Route{ - ID: nbroute.ID(fmt.Sprintf("network-id-%d", n)), - Description: "base route", - NetID: nbroute.NetID(fmt.Sprintf("network-id-%d", n)), - Network: netip.MustParsePrefix(netIP.String() + "/24"), - NetworkType: nbroute.IPv4Network, - Metric: 9999, - Masquerade: false, - Enabled: true, - Groups: []string{groupALL.ID}, - } - account.Routes[route.ID] = route - - group = &types.Group{ - ID: fmt.Sprintf("group-id-%d", n), - AccountID: account.Id, - Name: fmt.Sprintf("group-id-%d", n), - Issued: "api", - Peers: nil, - } - account.Groups[group.ID] = group - - nameserver := &nbdns.NameServerGroup{ - ID: fmt.Sprintf("nameserver-id-%d", n), - AccountID: account.Id, - Name: fmt.Sprintf("nameserver-id-%d", n), - Description: "", - NameServers: []nbdns.NameServer{{IP: netip.MustParseAddr(netIP.String()), NSType: nbdns.UDPNameServerType}}, - Groups: []string{group.ID}, - Primary: false, - Domains: nil, - Enabled: false, - SearchDomainsEnabled: false, - } - account.NameServerGroups[nameserver.ID] = nameserver - - setupKey, _ := types.GenerateDefaultSetupKey() - _, exists := account.SetupKeys[setupKey.Key] - if exists { - t.Errorf("setup key already exists") - } - account.SetupKeys[setupKey.Key] = setupKey - } - - err = store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - if len(store.GetAllAccounts(context.Background())) != 1 { - t.Errorf("expecting 1 Accounts to be stored after SaveAccount()") - } - - a, err := store.GetAccount(context.Background(), account.Id) - if a == nil { - t.Errorf("expecting Account to be stored after SaveAccount(): %v", err) - } - - if a != nil && len(a.Policies) != 1 { - t.Errorf("expecting Account to have one policy stored after SaveAccount(), got %d", len(a.Policies)) - } - - if a != nil && len(a.Policies[0].Rules) != 1 { - t.Errorf("expecting Account to have one policy rule stored after SaveAccount(), got %d", len(a.Policies[0].Rules)) - return - } - - if a != nil && len(a.Peers) != numPerAccount { - t.Errorf("expecting Account to have %d peers stored after SaveAccount(), got %d", - numPerAccount, len(a.Peers)) - return - } - - if a != nil && len(a.Users) != numPerAccount+1 { - t.Errorf("expecting Account to have %d users stored after SaveAccount(), got %d", - numPerAccount+1, len(a.Users)) - return - } - - if a != nil && len(a.Routes) != numPerAccount { - t.Errorf("expecting Account to have %d routes stored after SaveAccount(), got %d", - numPerAccount, len(a.Routes)) - return - } - - if a != nil && len(a.NameServerGroups) != numPerAccount { - t.Errorf("expecting Account to have %d NameServerGroups stored after SaveAccount(), got %d", - numPerAccount, len(a.NameServerGroups)) - return - } - - if a != nil && len(a.NameServerGroups) != numPerAccount { - t.Errorf("expecting Account to have %d NameServerGroups stored after SaveAccount(), got %d", - numPerAccount, len(a.NameServerGroups)) - return - } - - if a != nil && len(a.SetupKeys) != numPerAccount+1 { - t.Errorf("expecting Account to have %d SetupKeys stored after SaveAccount(), got %d", - numPerAccount+1, len(a.SetupKeys)) - return - } -} - -// sequentialIPv4 returns a unique IPv4 address for the given index, avoiding -// the random collisions that would otherwise violate the unique (account_id, ip) -// index when generating a large number of peers. -func sequentialIPv4(n int) net.IP { - b := make([]byte, 4) - binary.BigEndian.PutUint32(b, 0x0A000000+uint32(n)) - return net.IP(b) -} - -func Test_SaveAccount(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("The SQLite store is not properly supported by Windows yet") - } - - runTestForAllEngines(t, "", func(t *testing.T, store Store) { - account := newAccountWithId(context.Background(), "account_id", "testuser", "") - setupKey, _ := types.GenerateDefaultSetupKey() - account.SetupKeys[setupKey.Key] = setupKey - account.Peers["testpeer"] = &nbpeer.Peer{ - Key: "peerkey", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::1"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - - err := store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - account2 := newAccountWithId(context.Background(), "account_id2", "testuser2", "") - setupKey, _ = types.GenerateDefaultSetupKey() - account2.SetupKeys[setupKey.Key] = setupKey - account2.Peers["testpeer2"] = &nbpeer.Peer{ - Key: "peerkey2", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 2}), - IPv6: netip.MustParseAddr("fd00::2"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name 2", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - - err = store.SaveAccount(context.Background(), account2) - require.NoError(t, err) - - if len(store.GetAllAccounts(context.Background())) != 2 { - t.Errorf("expecting 2 Accounts to be stored after SaveAccount()") - } - - a, err := store.GetAccount(context.Background(), account.Id) - if a == nil { - t.Errorf("expecting Account to be stored after SaveAccount(): %v", err) - } - - if a != nil && len(a.Policies) != 1 { - t.Errorf("expecting Account to have one policy stored after SaveAccount(), got %d", len(a.Policies)) - } - - if a != nil && len(a.Policies[0].Rules) != 1 { - t.Errorf("expecting Account to have one policy rule stored after SaveAccount(), got %d", len(a.Policies[0].Rules)) - return - } - - if a, err := store.GetAccountByPeerPubKey(context.Background(), "peerkey"); a == nil { - t.Errorf("expecting PeerKeyID2AccountID index updated after SaveAccount(): %v", err) - } - - if a, err := store.GetAccountByUser(context.Background(), "testuser"); a == nil { - t.Errorf("expecting UserID2AccountID index updated after SaveAccount(): %v", err) - } - - if a, err := store.GetAccountByPeerID(context.Background(), "testpeer"); a == nil { - t.Errorf("expecting PeerID2AccountID index updated after SaveAccount(): %v", err) - } - - if a, err := store.GetAccountBySetupKey(context.Background(), setupKey.Key); a == nil { - t.Errorf("expecting SetupKeyID2AccountID index updated after SaveAccount(): %v", err) - } - }) -} - -func Test_AccountSettings_SaveAndRetrieve(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("The SQLite store is not properly supported by Windows yet") - } - - populateFields := testing_helpers.NewPopulateFields().WithCustomFieldSetter( - reflect.PointerTo(reflect.TypeOf(types.ExtraSettings{})), func(this *testing_helpers.PopulateFields, field reflect.Value) (int, error) { - es := types.ExtraSettings{} - reflectedEs := reflect.ValueOf(&es).Elem() - n, err := this.PopulateAll(reflectedEs) - if err != nil { - return n, err - } - field.Set(reflectedEs.Addr()) - return n, nil - }).WithCustomFieldSetter( - reflect.PointerTo(reflect.TypeOf(types.DashboardFeatures{})), func(this *testing_helpers.PopulateFields, field reflect.Value) (int, error) { - t := true - df := types.DashboardFeatures{AgentNetwork: &t} - reflectedDf := reflect.ValueOf(&df).Elem() - field.Set(reflectedDf.Addr()) - return 1, nil - }).WithSkippedTag("gorm", "-") - - runTestForAllEngines(t, "", func(t *testing.T, store Store) { - account := newAccountWithId(context.Background(), "account_id", "testuser", "") - setupKey, _ := types.GenerateDefaultSetupKey() - account.SetupKeys[setupKey.Key] = setupKey - - settings := types.Settings{} - numOfExportedFields, err := populateFields.PopulateAll(reflect.ValueOf(&settings).Elem()) - assert.NoError(t, err) - assert.Equal(t, 27, numOfExportedFields) - account.Settings = &settings - - err = store.SaveAccount(context.Background(), account) - assert.NoError(t, err) - - accountFromDb, err := store.GetAccount(context.Background(), account.Id) - assert.NoError(t, err) - assert.NotNil(t, accountFromDb) - assert.NotNil(t, accountFromDb.Settings) - - assert.True(t, reflect.DeepEqual(&settings, accountFromDb.Settings), "created settings and settings retrieved from the db should match") - }) -} - -func TestSqlite_DeleteAccount(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("The SQLite store is not properly supported by Windows yet") - } - - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - testUserID := "testuser" - user := types.NewAdminUser(testUserID) - user.PATs = map[string]*types.PersonalAccessToken{"testtoken": { - ID: "testtoken", - Name: "test token", - }} - - account := newAccountWithId(context.Background(), "account_id", testUserID, "") - setupKey, _ := types.GenerateDefaultSetupKey() - account.SetupKeys[setupKey.Key] = setupKey - account.Peers["testpeer"] = &nbpeer.Peer{ - Key: "peerkey", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::1"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - account.Users[testUserID] = user - account.Networks = []*networkTypes.Network{ - { - ID: "network_id", - AccountID: account.Id, - Name: "network name", - Description: "network description", - }, - } - account.NetworkRouters = []*routerTypes.NetworkRouter{ - { - ID: "router_id", - NetworkID: account.Networks[0].ID, - AccountID: account.Id, - PeerGroups: []string{"group_id"}, - Masquerade: true, - Metric: 1, - }, - } - account.NetworkResources = []*resourceTypes.NetworkResource{ - { - ID: "resource_id", - NetworkID: account.Networks[0].ID, - AccountID: account.Id, - Name: "Name", - Description: "Description", - Type: "Domain", - Address: "example.com", - }, - } - - account.Services = []*rpservice.Service{ - { - ID: "service_id", - AccountID: account.Id, - Name: "test service", - Domain: "svc.example.com", - Enabled: true, - Targets: []*rpservice.Target{ - { - AccountID: account.Id, - ServiceID: "service_id", - Host: "localhost", - Port: 8080, - Protocol: "http", - Enabled: true, - }, - }, - }, - } - - account.Domains = []*proxydomain.Domain{ - { - ID: "domain_id", - Domain: "custom.example.com", - AccountID: account.Id, - Validated: true, - }, - } - - err = store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - if len(store.GetAllAccounts(context.Background())) != 1 { - t.Errorf("expecting 1 Accounts to be stored after SaveAccount()") - } - - o, err := store.GetAccountOnboarding(context.Background(), account.Id) - require.NoError(t, err) - require.Equal(t, o.AccountID, account.Id) - - err = store.DeleteAccount(context.Background(), account) - require.NoError(t, err) - - _, err = store.GetAccountOnboarding(context.Background(), account.Id) - require.Error(t, err, "expecting error after removing DeleteAccount when getting onboarding") - - if len(store.GetAllAccounts(context.Background())) != 0 { - t.Errorf("expecting 0 Accounts to be stored after DeleteAccount()") - } - - _, err = store.GetAccountByPeerPubKey(context.Background(), "peerkey") - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer public key") - - _, err = store.GetAccountByUser(context.Background(), "testuser") - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by user") - - _, err = store.GetAccountByPeerID(context.Background(), "testpeer") - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer id") - - _, err = store.GetAccountBySetupKey(context.Background(), setupKey.Key) - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by setup key") - - _, err = store.GetAccount(context.Background(), account.Id) - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by id") - - for _, policy := range account.Policies { - var rules []*types.PolicyRule - err = store.(*SqlStore).db.Model(&types.PolicyRule{}).Find(&rules, "policy_id = ?", policy.ID).Error - require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for policy rules") - require.Len(t, rules, 0, "expecting no policy rules to be found after removing DeleteAccount") - - } - - for _, accountUser := range account.Users { - var pats []*types.PersonalAccessToken - err = store.(*SqlStore).db.Model(&types.PersonalAccessToken{}).Find(&pats, "user_id = ?", accountUser.Id).Error - require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for personal access token") - require.Len(t, pats, 0, "expecting no personal access token to be found after removing DeleteAccount") - - } - - for _, network := range account.Networks { - routers, err := store.GetNetworkRoutersByNetID(context.Background(), LockingStrengthNone, account.Id, network.ID) - require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for network routers") - require.Len(t, routers, 0, "expecting no network routers to be found after DeleteAccount") - - resources, err := store.GetNetworkResourcesByNetID(context.Background(), LockingStrengthNone, account.Id, network.ID) - require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for network resources") - require.Len(t, resources, 0, "expecting no network resources to be found after DeleteAccount") - } - - domains, err := store.ListCustomDomains(context.Background(), account.Id) - require.NoError(t, err, "expecting no error after DeleteAccount when searching for custom domains") - require.Len(t, domains, 0, "expecting no custom domains to be found after DeleteAccount") - - var services []*rpservice.Service - err = store.(*SqlStore).db.Model(&rpservice.Service{}).Find(&services, "account_id = ?", account.Id).Error - require.NoError(t, err, "expecting no error after DeleteAccount when searching for services") - require.Len(t, services, 0, "expecting no services to be found after DeleteAccount") - - var targets []*rpservice.Target - err = store.(*SqlStore).db.Model(&rpservice.Target{}).Find(&targets, "account_id = ?", account.Id).Error - require.NoError(t, err, "expecting no error after DeleteAccount when searching for service targets") - require.Len(t, targets, 0, "expecting no service targets to be found after DeleteAccount") -} - -func Test_GetAccount(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("The SQLite store is not properly supported by Windows yet") - } - - runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { - id := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - account, err := store.GetAccount(context.Background(), id) - require.NoError(t, err) - require.Equal(t, id, account.Id, "account id should match") - require.Equal(t, false, account.Onboarding.OnboardingFlowPending) - - id = "9439-34653001fc3b-bf1c8084-ba50-4ce7" - - account, err = store.GetAccount(context.Background(), id) - require.NoError(t, err) - require.Equal(t, id, account.Id, "account id should match") - require.Equal(t, true, account.Onboarding.OnboardingFlowPending) - - _, err = store.GetAccount(context.Background(), "non-existing-account") - assert.Error(t, err) - parsedErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") - - }) -} - -// TestSqlStore_GetPeerByIP_NotFound pins the not-found semantics the -// proxy's ValidateTunnelPeer relies on: a tunnel-IP that isn't in the -// account roster must surface as a NotFound error (not a generic -// Internal) so callers can distinguish an expected miss from a real -// store failure. A known IP still resolves. -func TestSqlStore_GetPeerByIP_NotFound(t *testing.T) { - runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { - const accountID = "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - peer, err := store.GetPeerByIP(context.Background(), LockingStrengthNone, accountID, net.ParseIP("192.168.0.0")) - require.NoError(t, err, "known tunnel IP must resolve") - require.NotNil(t, peer) - - _, err = store.GetPeerByIP(context.Background(), LockingStrengthNone, accountID, net.ParseIP("100.65.0.99")) - require.Error(t, err, "unknown tunnel IP must error") - parsedErr, ok := status.FromError(err) - require.True(t, ok, "error must be a status error") - require.Equal(t, status.NotFound, parsedErr.Type(), "tunnel-IP miss must be NotFound, not Internal") - }) -} - -func TestSqlStore_SavePeer(t *testing.T) { - populateFields := testing_helpers.NewPopulateFields() - - runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { - account, err := store.GetAccount(context.Background(), "bf1c8084-ba50-4ce7-9439-34653001fc3b") - require.NoError(t, err) - - metadata := nbpeer.PeerSystemMeta{} - reflectedMetadata := reflect.ValueOf(&metadata).Elem() - - numOfFields, err := populateFields.PopulateAll(reflectedMetadata) - assert.NoError(t, err) - assert.Equal(t, 33, numOfFields) - - // save status of non-existing peer - peer := &nbpeer.Peer{ - Key: "peerkey", - ID: "testpeer", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::1"), - Meta: metadata, //nbpeer.PeerSystemMeta{Hostname: "testingpeer"}, - Name: "peer name", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - CreatedAt: time.Now().UTC(), - } - ctx := context.Background() - err = store.SavePeer(ctx, account.Id, peer) - assert.Error(t, err) - parsedErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") - - // save new status of existing peer - account.Peers[peer.ID] = peer - - err = store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - updatedPeer := peer.Copy() - updatedPeer.Status.Connected = false - updatedPeer.Meta.Hostname = "updatedpeer" - - err = store.SavePeer(ctx, account.Id, updatedPeer) - require.NoError(t, err) - - account, err = store.GetAccount(context.Background(), account.Id) - require.NoError(t, err) - - actual := account.Peers[peer.ID] - assert.Equal(t, updatedPeer.Meta, actual.Meta) - assert.Equal(t, updatedPeer.Status.Connected, actual.Status.Connected) - assert.Equal(t, updatedPeer.Status.LoginExpired, actual.Status.LoginExpired) - assert.Equal(t, updatedPeer.Status.RequiresApproval, actual.Status.RequiresApproval) - assert.WithinDurationf(t, updatedPeer.Status.LastSeen, actual.Status.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") - }) -} - -func TestSqlStore_SavePeerStatus(t *testing.T) { - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - account, err := store.GetAccount(context.Background(), "bf1c8084-ba50-4ce7-9439-34653001fc3b") - require.NoError(t, err) - - // save status of non-existing peer - newStatus := nbpeer.PeerStatus{Connected: false, LastSeen: time.Now().UTC()} - err = store.SavePeerStatus(context.Background(), account.Id, "non-existing-peer", newStatus) - assert.Error(t, err) - parsedErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") - - // save new status of existing peer - account.Peers["testpeer"] = &nbpeer.Peer{ - Key: "peerkey", - ID: "testpeer", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::1"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - - err = store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - err = store.SavePeerStatus(context.Background(), account.Id, "testpeer", newStatus) - require.NoError(t, err) - - account, err = store.GetAccount(context.Background(), account.Id) - require.NoError(t, err) - - actual := account.Peers["testpeer"].Status - assert.Equal(t, newStatus.Connected, actual.Connected) - assert.Equal(t, newStatus.LoginExpired, actual.LoginExpired) - assert.Equal(t, newStatus.RequiresApproval, actual.RequiresApproval) - assert.WithinDurationf(t, newStatus.LastSeen, actual.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") - - newStatus.Connected = true - - err = store.SavePeerStatus(context.Background(), account.Id, "testpeer", newStatus) - require.NoError(t, err) - - account, err = store.GetAccount(context.Background(), account.Id) - require.NoError(t, err) - - actual = account.Peers["testpeer"].Status - assert.Equal(t, newStatus.Connected, actual.Connected) - assert.Equal(t, newStatus.LoginExpired, actual.LoginExpired) - assert.Equal(t, newStatus.RequiresApproval, actual.RequiresApproval) - assert.WithinDurationf(t, newStatus.LastSeen, actual.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") -} - -func Test_TestGetAccountByPrivateDomain(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("The SQLite store is not properly supported by Windows yet") - } - - runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { - existingDomain := "test.com" - - account, err := store.GetAccountByPrivateDomain(context.Background(), existingDomain) - require.NoError(t, err, "should found account") - require.Equal(t, existingDomain, account.Domain, "domains should match") - - _, err = store.GetAccountByPrivateDomain(context.Background(), "missing-domain.com") - require.Error(t, err, "should return error on domain lookup") - parsedErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") - }) -} - -func Test_GetTokenIDByHashedToken(t *testing.T) { - if runtime.GOOS == "windows" { - t.Skip("The SQLite store is not properly supported by Windows yet") - } - - runTestForAllEngines(t, "../testdata/store.sql", func(t *testing.T, store Store) { - hashed := "SoMeHaShEdToKeN" - id := "9dj38s35-63fb-11ec-90d6-0242ac120003" - - token, err := store.GetTokenIDByHashedToken(context.Background(), hashed) - require.NoError(t, err) - require.Equal(t, id, token) - - _, err = store.GetTokenIDByHashedToken(context.Background(), "non-existing-hash") - require.Error(t, err) - parsedErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, parsedErr.Type(), "should return not found error") - }) -} - func TestMigrate(t *testing.T) { if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { t.Skip("skip CI tests on darwin and windows") @@ -842,434 +183,6 @@ func TestPostgresql_NewStore(t *testing.T) { } } -func TestPostgresql_SaveAccount(t *testing.T) { - if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { - t.Skip("skip CI tests on darwin and windows") - } - - t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - account := newAccountWithId(context.Background(), "account_id", "testuser", "") - setupKey, _ := types.GenerateDefaultSetupKey() - account.SetupKeys[setupKey.Key] = setupKey - account.Peers["testpeer"] = &nbpeer.Peer{ - Key: "peerkey", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::1"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - - err = store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - account2 := newAccountWithId(context.Background(), "account_id2", "testuser2", "") - setupKey, _ = types.GenerateDefaultSetupKey() - account2.SetupKeys[setupKey.Key] = setupKey - account2.Peers["testpeer2"] = &nbpeer.Peer{ - Key: "peerkey2", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 2}), - IPv6: netip.MustParseAddr("fd00::2"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name 2", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - - err = store.SaveAccount(context.Background(), account2) - require.NoError(t, err) - - if len(store.GetAllAccounts(context.Background())) != 2 { - t.Errorf("expecting 2 Accounts to be stored after SaveAccount()") - } - - a, err := store.GetAccount(context.Background(), account.Id) - if a == nil { - t.Errorf("expecting Account to be stored after SaveAccount(): %v", err) - } - - if a != nil && len(a.Policies) != 1 { - t.Errorf("expecting Account to have one policy stored after SaveAccount(), got %d", len(a.Policies)) - } - - if a != nil && len(a.Policies[0].Rules) != 1 { - t.Errorf("expecting Account to have one policy rule stored after SaveAccount(), got %d", len(a.Policies[0].Rules)) - return - } - - if a, err := store.GetAccountByPeerPubKey(context.Background(), "peerkey"); a == nil { - t.Errorf("expecting PeerKeyID2AccountID index updated after SaveAccount(): %v", err) - } - - if a, err := store.GetAccountByUser(context.Background(), "testuser"); a == nil { - t.Errorf("expecting UserID2AccountID index updated after SaveAccount(): %v", err) - } - - if a, err := store.GetAccountByPeerID(context.Background(), "testpeer"); a == nil { - t.Errorf("expecting PeerID2AccountID index updated after SaveAccount(): %v", err) - } - - if a, err := store.GetAccountBySetupKey(context.Background(), setupKey.Key); a == nil { - t.Errorf("expecting SetupKeyID2AccountID index updated after SaveAccount(): %v", err) - } -} - -func TestPostgresql_DeleteAccount(t *testing.T) { - if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { - t.Skip("skip CI tests on darwin and windows") - } - - t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - testUserID := "testuser" - user := types.NewAdminUser(testUserID) - user.PATs = map[string]*types.PersonalAccessToken{"testtoken": { - ID: "testtoken", - Name: "test token", - }} - - account := newAccountWithId(context.Background(), "account_id", testUserID, "") - setupKey, _ := types.GenerateDefaultSetupKey() - account.SetupKeys[setupKey.Key] = setupKey - account.Peers["testpeer"] = &nbpeer.Peer{ - Key: "peerkey", - IP: netip.AddrFrom4([4]byte{127, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::1"), - Meta: nbpeer.PeerSystemMeta{}, - Name: "peer name", - Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now().UTC()}, - } - account.Users[testUserID] = user - - err = store.SaveAccount(context.Background(), account) - require.NoError(t, err) - - if len(store.GetAllAccounts(context.Background())) != 1 { - t.Errorf("expecting 1 Accounts to be stored after SaveAccount()") - } - - err = store.DeleteAccount(context.Background(), account) - require.NoError(t, err) - - if len(store.GetAllAccounts(context.Background())) != 0 { - t.Errorf("expecting 0 Accounts to be stored after DeleteAccount()") - } - - _, err = store.GetAccountByPeerPubKey(context.Background(), "peerkey") - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer public key") - - _, err = store.GetAccountByUser(context.Background(), "testuser") - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by user") - - _, err = store.GetAccountByPeerID(context.Background(), "testpeer") - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by peer id") - - _, err = store.GetAccountBySetupKey(context.Background(), setupKey.Key) - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by setup key") - - _, err = store.GetAccount(context.Background(), account.Id) - require.Error(t, err, "expecting error after removing DeleteAccount when getting account by id") - - for _, policy := range account.Policies { - var rules []*types.PolicyRule - err = store.(*SqlStore).db.Model(&types.PolicyRule{}).Find(&rules, "policy_id = ?", policy.ID).Error - require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for policy rules") - require.Len(t, rules, 0, "expecting no policy rules to be found after removing DeleteAccount") - - } - - for _, accountUser := range account.Users { - var pats []*types.PersonalAccessToken - err = store.(*SqlStore).db.Model(&types.PersonalAccessToken{}).Find(&pats, "user_id = ?", accountUser.Id).Error - require.NoError(t, err, "expecting no error after removing DeleteAccount when searching for personal access token") - require.Len(t, pats, 0, "expecting no personal access token to be found after removing DeleteAccount") - - } - -} - -func TestPostgresql_TestGetAccountByPrivateDomain(t *testing.T) { - if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { - t.Skip("skip CI tests on darwin and windows") - } - - t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - existingDomain := "test.com" - - account, err := store.GetAccountByPrivateDomain(context.Background(), existingDomain) - require.NoError(t, err, "should found account") - require.Equal(t, existingDomain, account.Domain, "domains should match") - - _, err = store.GetAccountByPrivateDomain(context.Background(), "missing-domain.com") - require.Error(t, err, "should return error on domain lookup") -} - -func TestPostgresql_GetTokenIDByHashedToken(t *testing.T) { - if (os.Getenv("CI") == "true" && runtime.GOOS == "darwin") || runtime.GOOS == "windows" { - t.Skip("skip CI tests on darwin and windows") - } - - t.Setenv("NETBIRD_STORE_ENGINE", string(types.PostgresStoreEngine)) - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - hashed := "SoMeHaShEdToKeN" - id := "9dj38s35-63fb-11ec-90d6-0242ac120003" - - token, err := store.GetTokenIDByHashedToken(context.Background(), hashed) - require.NoError(t, err) - require.Equal(t, id, token) -} - -func TestSqlite_GetTakenIPs(t *testing.T) { - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - defer cleanup() - if err != nil { - t.Fatal(err) - } - - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - _, err = store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - takenIPs, err := store.GetTakenIPs(context.Background(), LockingStrengthNone, existingAccountID) - require.NoError(t, err) - assert.Equal(t, []netip.Addr{}, takenIPs) - - peer1 := &nbpeer.Peer{ - ID: "peer1", - AccountID: existingAccountID, - Key: "key1", - DNSLabel: "peer1", - IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), - IPv6: netip.MustParseAddr("fd00::1:1:1:1"), - } - err = store.AddPeerToAccount(context.Background(), peer1) - require.NoError(t, err) - - takenIPs, err = store.GetTakenIPs(context.Background(), LockingStrengthNone, existingAccountID) - require.NoError(t, err) - ip1 := netip.AddrFrom4([4]byte{1, 1, 1, 1}) - assert.Equal(t, []netip.Addr{ip1}, takenIPs) - - peer2 := &nbpeer.Peer{ - ID: "peer1second", - AccountID: existingAccountID, - Key: "key2", - DNSLabel: "peer1-1", - IP: netip.AddrFrom4([4]byte{2, 2, 2, 2}), - IPv6: netip.MustParseAddr("fd00::2:2:2:2"), - } - err = store.AddPeerToAccount(context.Background(), peer2) - require.NoError(t, err) - - takenIPs, err = store.GetTakenIPs(context.Background(), LockingStrengthNone, existingAccountID) - require.NoError(t, err) - ip2 := netip.AddrFrom4([4]byte{2, 2, 2, 2}) - assert.Equal(t, []netip.Addr{ip1, ip2}, takenIPs) -} - -func TestSqlite_GetPeerLabelsInAccount(t *testing.T) { - runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - peerHostname := "peer1" - - _, err := store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - labels, err := store.GetPeerLabelsInAccount(context.Background(), LockingStrengthNone, existingAccountID, peerHostname) - require.NoError(t, err) - assert.Equal(t, []string{}, labels) - - peer1 := &nbpeer.Peer{ - ID: "peer1", - AccountID: existingAccountID, - Key: "key1", - DNSLabel: "peer1", - IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), - IPv6: netip.MustParseAddr("fd00::1:1:1:1"), - } - err = store.AddPeerToAccount(context.Background(), peer1) - require.NoError(t, err) - - labels, err = store.GetPeerLabelsInAccount(context.Background(), LockingStrengthNone, existingAccountID, peerHostname) - require.NoError(t, err) - assert.Equal(t, []string{"peer1"}, labels) - - peer2 := &nbpeer.Peer{ - ID: "peer1second", - AccountID: existingAccountID, - Key: "key2", - DNSLabel: "peer1-1", - IP: netip.AddrFrom4([4]byte{2, 2, 2, 2}), - IPv6: netip.MustParseAddr("fd00::2:2:2:2"), - } - err = store.AddPeerToAccount(context.Background(), peer2) - require.NoError(t, err) - - labels, err = store.GetPeerLabelsInAccount(context.Background(), LockingStrengthNone, existingAccountID, peerHostname) - require.NoError(t, err) - - expected := []string{"peer1", "peer1-1"} - sort.Strings(expected) - sort.Strings(labels) - assert.Equal(t, expected, labels) - }) -} - -func Test_AddPeerWithSameDnsLabel(t *testing.T) { - runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - _, err := store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - peer1 := &nbpeer.Peer{ - ID: "peer1", - AccountID: existingAccountID, - Key: "key1", - DNSLabel: "peer1.domain.test", - } - err = store.AddPeerToAccount(context.Background(), peer1) - require.NoError(t, err) - - peer2 := &nbpeer.Peer{ - ID: "peer1second", - AccountID: existingAccountID, - Key: "key2", - DNSLabel: "peer1.domain.test", - } - err = store.AddPeerToAccount(context.Background(), peer2) - require.Error(t, err) - }) -} - -func Test_AddPeerWithSameIP(t *testing.T) { - runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - _, err := store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - peer1 := &nbpeer.Peer{ - ID: "peer1", - AccountID: existingAccountID, - Key: "key1", - IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), - IPv6: netip.MustParseAddr("fd00::1:1:1:1"), - } - err = store.AddPeerToAccount(context.Background(), peer1) - require.NoError(t, err) - - peer2 := &nbpeer.Peer{ - ID: "peer1second", - AccountID: existingAccountID, - Key: "key2", - IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), - IPv6: netip.MustParseAddr("fd00::2:2:2:2"), - } - err = store.AddPeerToAccount(context.Background(), peer2) - require.Error(t, err) - }) -} - -func TestSqlite_GetAccountNetwork(t *testing.T) { - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - if err != nil { - t.Fatal(err) - } - - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - _, err = store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - network, err := store.GetAccountNetwork(context.Background(), LockingStrengthNone, existingAccountID) - require.NoError(t, err) - ip := net.IP{100, 64, 0, 0}.To16() - assert.Equal(t, ip, network.Net.IP) - assert.Equal(t, net.IPMask{255, 255, 0, 0}, network.Net.Mask) - assert.Equal(t, "", network.Dns) - assert.Equal(t, "af1c8024-ha40-4ce2-9418-34653101fc3c", network.Identifier) - assert.Equal(t, uint64(0), network.Serial) -} - -func TestSqlite_GetSetupKeyBySecret(t *testing.T) { - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - if err != nil { - t.Fatal(err) - } - - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - plainKey := "A2C8E62B-38F5-4553-B31E-DD66C696CEBB" - hashedKey := sha256.Sum256([]byte(plainKey)) - encodedHashedKey := b64.StdEncoding.EncodeToString(hashedKey[:]) - - _, err = store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - setupKey, err := store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) - require.NoError(t, err) - assert.Equal(t, encodedHashedKey, setupKey.Key) - assert.Equal(t, types.HiddenKey(plainKey, 4), setupKey.KeySecret) - assert.Equal(t, "bf1c8084-ba50-4ce7-9439-34653001fc3b", setupKey.AccountID) - assert.Equal(t, "Default key", setupKey.Name) -} - -func TestSqlite_incrementSetupKeyUsage(t *testing.T) { - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - if err != nil { - t.Fatal(err) - } - - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - plainKey := "A2C8E62B-38F5-4553-B31E-DD66C696CEBB" - hashedKey := sha256.Sum256([]byte(plainKey)) - encodedHashedKey := b64.StdEncoding.EncodeToString(hashedKey[:]) - - _, err = store.GetAccount(context.Background(), existingAccountID) - require.NoError(t, err) - - setupKey, err := store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) - require.NoError(t, err) - assert.Equal(t, 0, setupKey.UsedTimes) - - err = store.IncrementSetupKeyUsage(context.Background(), setupKey.Id) - require.NoError(t, err) - - setupKey, err = store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) - require.NoError(t, err) - assert.Equal(t, 1, setupKey.UsedTimes) - - err = store.IncrementSetupKeyUsage(context.Background(), setupKey.Id) - require.NoError(t, err) - - setupKey, err = store.GetSetupKeyBySecret(context.Background(), LockingStrengthNone, encodedHashedKey) - require.NoError(t, err) - assert.Equal(t, 2, setupKey.UsedTimes) -} - func TestSqlite_CreateAndGetObjectInTransaction(t *testing.T) { t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) @@ -1302,939 +215,6 @@ func TestSqlite_CreateAndGetObjectInTransaction(t *testing.T) { assert.NoError(t, err) } -func TestSqlStore_SaveAccountPersistsAgentNetworkOnly(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - account, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.False(t, account.Settings.AgentNetworkOnly, "setting should default to false") - - account.Settings.AgentNetworkOnly = true - require.NoError(t, store.SaveAccount(context.Background(), account)) - - reloaded, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.True(t, reloaded.Settings.AgentNetworkOnly, "setting should survive a save/load round-trip") - - reloaded.Settings.AgentNetworkOnly = false - require.NoError(t, store.SaveAccount(context.Background(), reloaded)) - - disabled, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.False(t, disabled.Settings.AgentNetworkOnly, "disabling should persist") -} - -func TestSqlStore_SaveAccountPersistsDashboardFeatures(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - account, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.Nil(t, account.Settings.DashboardFeatures, "dashboard features should default to unset") - - agentNetwork := true - account.Settings.DashboardFeatures = &types.DashboardFeatures{AgentNetwork: &agentNetwork} - require.NoError(t, store.SaveAccount(context.Background(), account)) - - reloaded, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.NotNil(t, reloaded.Settings.DashboardFeatures, "dashboard features should survive a save/load round-trip") - require.NotNil(t, reloaded.Settings.DashboardFeatures.AgentNetwork, "agent network flag should be set") - require.True(t, *reloaded.Settings.DashboardFeatures.AgentNetwork, "agent network flag should persist as true") - - disabled := false - reloaded.Settings.DashboardFeatures = &types.DashboardFeatures{AgentNetwork: &disabled} - require.NoError(t, store.SaveAccount(context.Background(), reloaded)) - - reloadedDisabled, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.NotNil(t, reloadedDisabled.Settings.DashboardFeatures.AgentNetwork, "agent network flag should remain set") - require.False(t, *reloadedDisabled.Settings.DashboardFeatures.AgentNetwork, "explicit false should persist") -} - -func TestSqlStore_GetAccountUsers(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - if err != nil { - t.Fatal(err) - } - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - account, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - users, err := store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Len(t, users, len(account.Users)) -} - -func TestSqlStore_UpdateAccountDomainAttributes(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - if err != nil { - t.Fatal(err) - } - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - t.Run("Should update attributes with public domain", func(t *testing.T) { - require.NoError(t, err) - domain := "example.com" - category := "public" - IsDomainPrimaryAccount := false - err = store.UpdateAccountDomainAttributes(context.Background(), accountID, domain, category, IsDomainPrimaryAccount) - require.NoError(t, err) - account, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.Equal(t, domain, account.Domain) - require.Equal(t, category, account.DomainCategory) - require.Equal(t, IsDomainPrimaryAccount, account.IsDomainPrimaryAccount) - }) - - t.Run("Should update attributes with private domain", func(t *testing.T) { - require.NoError(t, err) - domain := "test.com" - category := "private" - IsDomainPrimaryAccount := true - err = store.UpdateAccountDomainAttributes(context.Background(), accountID, domain, category, IsDomainPrimaryAccount) - require.NoError(t, err) - account, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - require.Equal(t, domain, account.Domain) - require.Equal(t, category, account.DomainCategory) - require.Equal(t, IsDomainPrimaryAccount, account.IsDomainPrimaryAccount) - }) - - t.Run("Should fail when account does not exist", func(t *testing.T) { - require.NoError(t, err) - domain := "test.com" - category := "private" - IsDomainPrimaryAccount := true - err = store.UpdateAccountDomainAttributes(context.Background(), "non-existing-account-id", domain, category, IsDomainPrimaryAccount) - require.Error(t, err) - }) - -} - -func TestSqlite_GetGroupByName(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - if err != nil { - t.Fatal(err) - } - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - group, err := store.GetGroupByName(context.Background(), LockingStrengthNone, accountID, "All") - require.NoError(t, err) - require.True(t, group.IsGroupAll()) -} - -func Test_DeleteSetupKeySuccessfully(t *testing.T) { - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - setupKeyID := "A2C8E62B-38F5-4553-B31E-DD66C696CEBB" - - err = store.DeleteSetupKey(context.Background(), accountID, setupKeyID) - require.NoError(t, err) - - _, err = store.GetSetupKeyByID(context.Background(), LockingStrengthNone, setupKeyID, accountID) - require.Error(t, err) -} - -func Test_DeleteSetupKeyFailsForNonExistingKey(t *testing.T) { - t.Setenv("NETBIRD_STORE_ENGINE", string(types.SqliteStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - nonExistingKeyID := "non-existing-key-id" - - err = store.DeleteSetupKey(context.Background(), accountID, nonExistingKeyID) - require.Error(t, err) -} - -func TestSqlStore_GetGroupsByIDs(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - groupIDs []string - expectedCount int - }{ - { - name: "retrieve existing groups by existing IDs", - groupIDs: []string{"cfefqs706sqkneg59g4g", "cfefqs706sqkneg59g3g"}, - expectedCount: 2, - }, - { - name: "empty group IDs list", - groupIDs: []string{}, - expectedCount: 0, - }, - { - name: "non-existing group IDs", - groupIDs: []string{"nonexistent1", "nonexistent2"}, - expectedCount: 0, - }, - { - name: "mixed existing and non-existing group IDs", - groupIDs: []string{"cfefqs706sqkneg59g4g", "nonexistent"}, - expectedCount: 1, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - groups, err := store.GetGroupsByIDs(context.Background(), LockingStrengthNone, accountID, tt.groupIDs) - require.NoError(t, err) - require.Len(t, groups, tt.expectedCount) - }) - } -} - -func TestSqlStore_CreateGroup(t *testing.T) { - if os.Getenv("CI") == "true" { - t.Log("Skipping MySQL test on CI") - } - t.Setenv("NETBIRD_STORE_ENGINE", string(types.MysqlStoreEngine)) - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - group := &types.Group{ - ID: "group-id", - AccountID: accountID, - Issued: "api", - Peers: []string{}, - Resources: []types.Resource{}, - GroupPeers: []types.GroupPeer{}, - } - err = store.CreateGroup(context.Background(), group) - require.NoError(t, err) - - savedGroup, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, "group-id") - require.NoError(t, err) - require.Equal(t, savedGroup, group) -} - -func TestSqlStore_CreateUpdateGroups(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - groups := []*types.Group{ - { - ID: "group-1", - AccountID: accountID, - Issued: "api", - Peers: []string{}, - Resources: []types.Resource{}, - GroupPeers: []types.GroupPeer{}, - }, - { - ID: "group-2", - AccountID: accountID, - Issued: "integration", - Peers: []string{}, - Resources: []types.Resource{}, - GroupPeers: []types.GroupPeer{}, - }, - } - err = store.CreateGroups(context.Background(), accountID, groups) - require.NoError(t, err) - - groups[1].Peers = []string{} - err = store.UpdateGroups(context.Background(), accountID, groups) - require.NoError(t, err) - - group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groups[1].ID) - require.NoError(t, err) - require.Equal(t, groups[1], group) -} - -func TestSqlStore_DeleteGroup(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - groupID string - expectError bool - }{ - { - name: "delete existing group", - groupID: "cfefqs706sqkneg59g4g", - expectError: false, - }, - { - name: "delete non-existing group", - groupID: "non-existing-group-id", - expectError: true, - }, - { - name: "delete with empty group ID", - groupID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - err := store.DeleteGroup(context.Background(), accountID, tt.groupID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - } else { - require.NoError(t, err) - - group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, tt.groupID) - require.Error(t, err) - require.Nil(t, group) - } - }) - } -} - -func TestSqlStore_DeleteGroups(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - groupIDs []string - expectError bool - }{ - { - name: "delete multiple existing groups", - groupIDs: []string{"cfefqs706sqkneg59g4g", "cfefqs706sqkneg59g3g"}, - expectError: false, - }, - { - name: "delete non-existing groups", - groupIDs: []string{"non-existing-id-1", "non-existing-id-2"}, - expectError: false, - }, - { - name: "delete with empty group IDs list", - groupIDs: []string{}, - expectError: false, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - err := store.DeleteGroups(context.Background(), accountID, tt.groupIDs) - if tt.expectError { - require.Error(t, err) - } else { - require.NoError(t, err) - - for _, groupID := range tt.groupIDs { - group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.Error(t, err) - require.Nil(t, group) - } - } - }) - } -} - -func TestSqlStore_GetPeerByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - peerID string - expectError bool - }{ - { - name: "retrieve existing peer", - peerID: "cfefqs706sqkneg59g4g", - expectError: false, - }, - { - name: "retrieve non-existing peer", - peerID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty peer ID", - peerID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peer, err := store.GetPeerByID(context.Background(), LockingStrengthNone, accountID, tt.peerID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, peer) - } else { - require.NoError(t, err) - require.NotNil(t, peer) - require.Equal(t, tt.peerID, peer.ID) - } - }) - } -} - -func TestSqlStore_GetPeersByIDs(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - peerIDs []string - expectedCount int - }{ - { - name: "retrieve existing peers by existing IDs", - peerIDs: []string{"cfefqs706sqkneg59g4g", "cfeg6sf06sqkneg59g50"}, - expectedCount: 2, - }, - { - name: "empty peer IDs list", - peerIDs: []string{}, - expectedCount: 0, - }, - { - name: "non-existing peer IDs", - peerIDs: []string{"nonexistent1", "nonexistent2"}, - expectedCount: 0, - }, - { - name: "mixed existing and non-existing peer IDs", - peerIDs: []string{"cfeg6sf06sqkneg59g50", "nonexistent"}, - expectedCount: 1, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetPeersByIDs(context.Background(), LockingStrengthNone, accountID, tt.peerIDs) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - }) - } -} - -func TestSqlStore_GetPostureChecksByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - postureChecksID string - expectError bool - }{ - { - name: "retrieve existing posture checks", - postureChecksID: "csplshq7qv948l48f7t0", - expectError: false, - }, - { - name: "retrieve non-existing posture checks", - postureChecksID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty posture checks ID", - postureChecksID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - postureChecks, err := store.GetPostureChecksByID(context.Background(), LockingStrengthNone, accountID, tt.postureChecksID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, postureChecks) - } else { - require.NoError(t, err) - require.NotNil(t, postureChecks) - require.Equal(t, tt.postureChecksID, postureChecks.ID) - } - }) - } -} - -func TestSqlStore_GetPostureChecksByIDs(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - postureCheckIDs []string - expectedCount int - }{ - { - name: "retrieve existing posture checks by existing IDs", - postureCheckIDs: []string{"csplshq7qv948l48f7t0", "cspnllq7qv95uq1r4k90"}, - expectedCount: 2, - }, - { - name: "empty posture check IDs list", - postureCheckIDs: []string{}, - expectedCount: 0, - }, - { - name: "non-existing posture check IDs", - postureCheckIDs: []string{"nonexistent1", "nonexistent2"}, - expectedCount: 0, - }, - { - name: "mixed existing and non-existing posture check IDs", - postureCheckIDs: []string{"cspnllq7qv95uq1r4k90", "nonexistent"}, - expectedCount: 1, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - groups, err := store.GetPostureChecksByIDs(context.Background(), LockingStrengthNone, accountID, tt.postureCheckIDs) - require.NoError(t, err) - require.Len(t, groups, tt.expectedCount) - }) - } -} - -func TestSqlStore_SavePostureChecks(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - postureChecks := &posture.Checks{ - ID: "posture-checks-id", - AccountID: accountID, - Checks: posture.ChecksDefinition{ - NBVersionCheck: &posture.NBVersionCheck{ - MinVersion: "0.31.0", - }, - OSVersionCheck: &posture.OSVersionCheck{ - Ios: &posture.MinVersionCheck{ - MinVersion: "13.0.1", - }, - Linux: &posture.MinKernelVersionCheck{ - MinKernelVersion: "5.3.3-dev", - }, - }, - GeoLocationCheck: &posture.GeoLocationCheck{ - Locations: []posture.Location{ - { - CountryCode: "DE", - CityName: "Berlin", - }, - }, - Action: posture.CheckActionAllow, - }, - }, - } - err = store.SavePostureChecks(context.Background(), postureChecks) - require.NoError(t, err) - - savePostureChecks, err := store.GetPostureChecksByID(context.Background(), LockingStrengthNone, accountID, "posture-checks-id") - require.NoError(t, err) - require.Equal(t, savePostureChecks, postureChecks) -} - -func TestSqlStore_DeletePostureChecks(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - postureChecksID string - expectError bool - }{ - { - name: "delete existing posture checks", - postureChecksID: "csplshq7qv948l48f7t0", - expectError: false, - }, - { - name: "delete non-existing posture checks", - postureChecksID: "non-existing-posture-checks-id", - expectError: true, - }, - { - name: "delete with empty posture checks ID", - postureChecksID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - err = store.DeletePostureChecks(context.Background(), accountID, tt.postureChecksID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - } else { - require.NoError(t, err) - group, err := store.GetPostureChecksByID(context.Background(), LockingStrengthNone, accountID, tt.postureChecksID) - require.Error(t, err) - require.Nil(t, group) - } - }) - } -} - -func TestSqlStore_GetPolicyByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - policyID string - expectError bool - }{ - { - name: "retrieve existing policy", - policyID: "cs1tnh0hhcjnqoiuebf0", - expectError: false, - }, - { - name: "retrieve non-existing policy checks", - policyID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty policy ID", - policyID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, tt.policyID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, policy) - } else { - require.NoError(t, err) - require.NotNil(t, policy) - require.Equal(t, tt.policyID, policy.ID) - } - }) - } -} - -func TestSqlStore_CreatePolicy(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - policy := &types.Policy{ - ID: "policy-id", - AccountID: accountID, - Enabled: true, - Rules: []*types.PolicyRule{ - { - Enabled: true, - Sources: []string{"groupA"}, - Destinations: []string{"groupC"}, - Bidirectional: true, - Action: types.PolicyTrafficActionAccept, - }, - }, - } - err = store.CreatePolicy(context.Background(), policy) - require.NoError(t, err) - - savePolicy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policy.ID) - require.NoError(t, err) - require.Equal(t, savePolicy, policy) - -} - -func TestSqlStore_SavePolicy(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - policyID := "cs1tnh0hhcjnqoiuebf0" - - policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policyID) - require.NoError(t, err) - - policy.Enabled = false - policy.Description = "policy" - policy.Rules[0].Sources = []string{"group"} - policy.Rules[0].Ports = []string{"80", "443"} - err = store.SavePolicy(context.Background(), policy) - require.NoError(t, err) - - savePolicy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policy.ID) - require.NoError(t, err) - require.Equal(t, savePolicy, policy) -} - -func TestSqlStore_DeletePolicy(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - policyID := "cs1tnh0hhcjnqoiuebf0" - - err = store.DeletePolicy(context.Background(), accountID, policyID) - require.NoError(t, err) - - policy, err := store.GetPolicyByID(context.Background(), LockingStrengthNone, accountID, policyID) - require.Error(t, err) - require.Nil(t, policy) -} - -func TestSqlStore_GetDNSSettings(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectError bool - }{ - { - name: "retrieve existing account dns settings", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectError: false, - }, - { - name: "retrieve non-existing account dns settings", - accountID: "non-existing", - expectError: true, - }, - { - name: "retrieve dns settings with empty account ID", - accountID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - dnsSettings, err := store.GetAccountDNSSettings(context.Background(), LockingStrengthNone, tt.accountID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, dnsSettings) - } else { - require.NoError(t, err) - require.NotNil(t, dnsSettings) - } - }) - } -} - -func TestSqlStore_SaveDNSSettings(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - dnsSettings, err := store.GetAccountDNSSettings(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - - dnsSettings.DisabledManagementGroups = []string{"groupA", "groupB"} - err = store.SaveDNSSettings(context.Background(), accountID, dnsSettings) - require.NoError(t, err) - - saveDNSSettings, err := store.GetAccountDNSSettings(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Equal(t, saveDNSSettings, dnsSettings) -} - -func TestSqlStore_GetAccountNameServerGroups(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectedCount int - }{ - { - name: "retrieve name server groups by existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectedCount: 1, - }, - { - name: "non-existing account ID", - accountID: "nonexistent", - expectedCount: 0, - }, - { - name: "empty account ID", - accountID: "", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetAccountNameServerGroups(context.Background(), LockingStrengthNone, tt.accountID) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - }) - } - -} - -func TestSqlStore_GetNameServerByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - nsGroupID string - expectError bool - }{ - { - name: "retrieve existing nameserver group", - nsGroupID: "csqdelq7qv97ncu7d9t0", - expectError: false, - }, - { - name: "retrieve non-existing nameserver group", - nsGroupID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty nameserver group ID", - nsGroupID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - nsGroup, err := store.GetNameServerGroupByID(context.Background(), LockingStrengthNone, accountID, tt.nsGroupID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, nsGroup) - } else { - require.NoError(t, err) - require.NotNil(t, nsGroup) - require.Equal(t, tt.nsGroupID, nsGroup.ID) - } - }) - } -} - -func TestSqlStore_SaveNameServerGroup(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - nsGroup := &nbdns.NameServerGroup{ - ID: "ns-group-id", - AccountID: accountID, - Name: "NS Group", - NameServers: []nbdns.NameServer{ - { - IP: netip.MustParseAddr("8.8.8.8"), - NSType: 1, - Port: 53, - }, - }, - Groups: []string{"groupA"}, - Primary: true, - Enabled: true, - SearchDomainsEnabled: false, - } - - err = store.SaveNameServerGroup(context.Background(), nsGroup) - require.NoError(t, err) - - saveNSGroup, err := store.GetNameServerGroupByID(context.Background(), LockingStrengthNone, accountID, nsGroup.ID) - require.NoError(t, err) - require.Equal(t, saveNSGroup, nsGroup) -} - -func TestSqlStore_DeleteNameServerGroup(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - nsGroupID := "csqdelq7qv97ncu7d9t0" - - err = store.DeleteNameServerGroup(context.Background(), accountID, nsGroupID) - require.NoError(t, err) - - nsGroup, err := store.GetNameServerGroupByID(context.Background(), LockingStrengthNone, accountID, nsGroupID) - require.Error(t, err) - require.Nil(t, nsGroup) -} - // newAccountWithId creates a new Account with a default SetupKey (doesn't store in a Store) and provided id func newAccountWithId(ctx context.Context, accountID, userID, domain string) *types.Account { log.WithContext(ctx).Debugf("creating new account") @@ -2285,805 +265,6 @@ func newAccountWithId(ctx context.Context, accountID, userID, domain string) *ty return acc } -func TestSqlStore_GetAccountNetworks(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectedCount int - }{ - { - name: "retrieve networks by existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectedCount: 1, - }, - - { - name: "retrieve networks by non-existing account ID", - accountID: "non-existent", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - networks, err := store.GetAccountNetworks(context.Background(), LockingStrengthNone, tt.accountID) - require.NoError(t, err) - require.Len(t, networks, tt.expectedCount) - }) - } -} - -func TestSqlStore_GetNetworkByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - networkID string - expectError bool - }{ - { - name: "retrieve existing network ID", - networkID: "ct286bi7qv930dsrrug0", - expectError: false, - }, - { - name: "retrieve non-existing network ID", - networkID: "non-existing", - expectError: true, - }, - { - name: "retrieve network with empty ID", - networkID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - network, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, tt.networkID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, network) - } else { - require.NoError(t, err) - require.NotNil(t, network) - require.Equal(t, tt.networkID, network.ID) - } - }) - } -} - -func TestSqlStore_SaveNetwork(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - network := &networkTypes.Network{ - ID: "net-id", - AccountID: accountID, - Name: "net", - } - - err = store.SaveNetwork(context.Background(), network) - require.NoError(t, err) - - savedNet, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, network.ID) - require.NoError(t, err) - require.Equal(t, network, savedNet) -} - -func TestSqlStore_DeleteNetwork(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - networkID := "ct286bi7qv930dsrrug0" - - err = store.DeleteNetwork(context.Background(), accountID, networkID) - require.NoError(t, err) - - network, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, networkID) - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, sErr.Type()) - require.Nil(t, network) -} - -func TestSqlStore_GetNetworkRoutersByNetID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - networkID string - expectedCount int - }{ - { - name: "retrieve routers by existing network ID", - networkID: "ct286bi7qv930dsrrug0", - expectedCount: 1, - }, - { - name: "retrieve routers by non-existing network ID", - networkID: "non-existent", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - routers, err := store.GetNetworkRoutersByNetID(context.Background(), LockingStrengthNone, accountID, tt.networkID) - require.NoError(t, err) - require.Len(t, routers, tt.expectedCount) - }) - } -} - -func TestSqlStore_GetNetworkRouterByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - networkRouterID string - expectError bool - }{ - { - name: "retrieve existing network router ID", - networkRouterID: "ctc20ji7qv9ck2sebc80", - expectError: false, - }, - { - name: "retrieve non-existing network router ID", - networkRouterID: "non-existing", - expectError: true, - }, - { - name: "retrieve network with empty router ID", - networkRouterID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - networkRouter, err := store.GetNetworkRouterByID(context.Background(), LockingStrengthNone, accountID, tt.networkRouterID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, networkRouter) - } else { - require.NoError(t, err) - require.NotNil(t, networkRouter) - require.Equal(t, tt.networkRouterID, networkRouter.ID) - } - }) - } -} - -func TestSqlStore_CreateNetworkRouter(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - networkID := "ct286bi7qv930dsrrug0" - - netRouter, err := routerTypes.NewNetworkRouter(accountID, networkID, "", []string{"net-router-grp"}, true, 0, true) - require.NoError(t, err) - - err = store.CreateNetworkRouter(context.Background(), netRouter) - require.NoError(t, err) - - savedNetRouter, err := store.GetNetworkRouterByID(context.Background(), LockingStrengthNone, accountID, netRouter.ID) - require.NoError(t, err) - require.Equal(t, netRouter, savedNetRouter) -} - -func TestSqlStore_UpdateNetworkRouter(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - networkID := "ct286bi7qv930dsrrug0" - routerID := "ctc20ji7qv9ck2sebc80" - - netRouter := &routerTypes.NetworkRouter{ - ID: routerID, - AccountID: accountID, - NetworkID: networkID, - Peer: "", - PeerGroups: []string{"net-router-grp"}, - Masquerade: true, - Metric: 42, - Enabled: true, - } - - err = store.UpdateNetworkRouter(context.Background(), netRouter) - require.NoError(t, err) - - savedNetRouter, err := store.GetNetworkRouterByID(context.Background(), LockingStrengthNone, accountID, routerID) - require.NoError(t, err) - require.Equal(t, netRouter, savedNetRouter) - - // Updating a router under a different account must not match any row. - netRouter.AccountID = "non-existent-account" - err = store.UpdateNetworkRouter(context.Background(), netRouter) - require.Error(t, err) -} - -func TestSqlStore_DeleteNetworkRouter(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - netRouterID := "ctc20ji7qv9ck2sebc80" - - err = store.DeleteNetworkRouter(context.Background(), accountID, netRouterID) - require.NoError(t, err) - - netRouter, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, netRouterID) - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, sErr.Type()) - require.Nil(t, netRouter) -} - -func TestSqlStore_GetNetworkResourcesByNetID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - tests := []struct { - name string - networkID string - expectedCount int - }{ - { - name: "retrieve resources by existing network ID", - networkID: "ct286bi7qv930dsrrug0", - expectedCount: 1, - }, - { - name: "retrieve resources by non-existing network ID", - networkID: "non-existent", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - netResources, err := store.GetNetworkResourcesByNetID(context.Background(), LockingStrengthNone, accountID, tt.networkID) - require.NoError(t, err) - require.Len(t, netResources, tt.expectedCount) - }) - } -} - -func TestSqlStore_GetNetworkResourceByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - netResourceID string - expectError bool - }{ - { - name: "retrieve existing network resource ID", - netResourceID: "ctc4nci7qv9061u6ilfg", - expectError: false, - }, - { - name: "retrieve non-existing network resource ID", - netResourceID: "non-existing", - expectError: true, - }, - { - name: "retrieve network with empty resource ID", - netResourceID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - netResource, err := store.GetNetworkResourceByID(context.Background(), LockingStrengthNone, accountID, tt.netResourceID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, netResource) - } else { - require.NoError(t, err) - require.NotNil(t, netResource) - require.Equal(t, tt.netResourceID, netResource.ID) - } - }) - } -} - -func TestSqlStore_SaveNetworkResource(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - networkID := "ct286bi7qv930dsrrug0" - - netResource, err := resourceTypes.NewNetworkResource(accountID, networkID, "resource-name", "", "example.com", []string{}, true) - require.NoError(t, err) - - err = store.SaveNetworkResource(context.Background(), netResource) - require.NoError(t, err) - - savedNetResource, err := store.GetNetworkResourceByID(context.Background(), LockingStrengthNone, accountID, netResource.ID) - require.NoError(t, err) - require.Equal(t, netResource.ID, savedNetResource.ID) - require.Equal(t, netResource.Name, savedNetResource.Name) - require.Equal(t, netResource.NetworkID, savedNetResource.NetworkID) - require.Equal(t, netResource.Type, resourceTypes.NetworkResourceType("domain")) - require.Equal(t, netResource.Domain, "example.com") - require.Equal(t, netResource.AccountID, savedNetResource.AccountID) - require.Equal(t, netResource.Prefix, netip.Prefix{}) -} - -func TestSqlStore_DeleteNetworkResource(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - netResourceID := "ctc4nci7qv9061u6ilfg" - - err = store.DeleteNetworkResource(context.Background(), accountID, netResourceID) - require.NoError(t, err) - - netResource, err := store.GetNetworkByID(context.Background(), LockingStrengthNone, accountID, netResourceID) - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, status.NotFound, sErr.Type()) - require.Nil(t, netResource) -} - -func TestSqlStore_AddAndRemoveResourceFromGroup(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - require.NoError(t, err) - t.Cleanup(cleanup) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - resourceId := "ctc4nci7qv9061u6ilfg" - groupID := "cs1tnh0hhcjnqoiuebeg" - - res := &types.Resource{ - ID: resourceId, - Type: "host", - } - err = store.AddResourceToGroup(context.Background(), accountID, groupID, res) - require.NoError(t, err) - - group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.NoError(t, err) - require.Contains(t, group.Resources, *res) - - groups, err := store.GetResourceGroups(context.Background(), LockingStrengthNone, accountID, resourceId) - require.NoError(t, err) - require.Len(t, groups, 1) - - err = store.RemoveResourceFromGroup(context.Background(), accountID, groupID, res.ID) - require.NoError(t, err) - - group, err = store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.NoError(t, err) - require.NotContains(t, group.Resources, *res) -} - -func TestSqlStore_AddPeerToGroup(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - peerID := "cfefqs706sqkneg59g4g" - groupID := "cfefqs706sqkneg59g4h" - - group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.NoError(t, err, "failed to get group") - require.Len(t, group.Peers, 0, "group should have 0 peers") - - err = store.AddPeerToGroup(context.Background(), accountID, peerID, groupID) - require.NoError(t, err, "failed to add peer to group") - - group, err = store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.NoError(t, err, "failed to get group") - require.Len(t, group.Peers, 1, "group should have 1 peers") - require.Contains(t, group.Peers, peerID) -} - -func TestSqlStore_AddPeerToAllGroup(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - groupID := "cfefqs706sqkneg59g3g" - - peer := &nbpeer.Peer{ - ID: "peer1", - AccountID: accountID, - DNSLabel: "peer1.domain.test", - } - - group, err := store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.NoError(t, err, "failed to get group") - require.Len(t, group.Peers, 2, "group should have 2 peers") - require.NotContains(t, group.Peers, peer.ID) - - err = store.AddPeerToAccount(context.Background(), peer) - require.NoError(t, err, "failed to add peer to account") - - err = store.AddPeerToAllGroup(context.Background(), accountID, peer.ID) - require.NoError(t, err, "failed to add peer to all group") - - group, err = store.GetGroupByID(context.Background(), LockingStrengthNone, accountID, groupID) - require.NoError(t, err, "failed to get group") - require.Len(t, group.Peers, 3, "group should have peers") - require.Contains(t, group.Peers, peer.ID) -} - -func TestSqlStore_AddPeerToAccount(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - peer := &nbpeer.Peer{ - ID: "peer1", - AccountID: accountID, - Key: "key", - IP: netip.AddrFrom4([4]byte{1, 1, 1, 1}), - IPv6: netip.MustParseAddr("fd00::1:1:1:1"), - Meta: nbpeer.PeerSystemMeta{ - Hostname: "hostname", - GoOS: "linux", - Kernel: "Linux", - Core: "21.04", - Platform: "x86_64", - OS: "Ubuntu", - WtVersion: "development", - UIVersion: "development", - }, - Name: "peer.test", - DNSLabel: "peer", - Status: &nbpeer.PeerStatus{ - LastSeen: time.Now().UTC(), - Connected: true, - LoginExpired: false, - RequiresApproval: false, - }, - SSHKey: "ssh-key", - SSHEnabled: false, - LoginExpirationEnabled: true, - InactivityExpirationEnabled: false, - LastLogin: util.ToPtr(time.Now().UTC()), - CreatedAt: time.Now().UTC(), - Ephemeral: true, - } - err = store.AddPeerToAccount(context.Background(), peer) - require.NoError(t, err, "failed to add peer to account") - - storedPeer, err := store.GetPeerByID(context.Background(), LockingStrengthNone, accountID, peer.ID) - require.NoError(t, err, "failed to get peer") - - assert.Equal(t, peer.ID, storedPeer.ID) - assert.Equal(t, peer.AccountID, storedPeer.AccountID) - assert.Equal(t, peer.Key, storedPeer.Key) - assert.Equal(t, peer.IP.String(), storedPeer.IP.String()) - assert.Equal(t, peer.Meta, storedPeer.Meta) - assert.Equal(t, peer.Name, storedPeer.Name) - assert.Equal(t, peer.DNSLabel, storedPeer.DNSLabel) - assert.Equal(t, peer.SSHKey, storedPeer.SSHKey) - assert.Equal(t, peer.SSHEnabled, storedPeer.SSHEnabled) - assert.Equal(t, peer.LoginExpirationEnabled, storedPeer.LoginExpirationEnabled) - assert.Equal(t, peer.InactivityExpirationEnabled, storedPeer.InactivityExpirationEnabled) - assert.WithinDurationf(t, peer.GetLastLogin(), storedPeer.GetLastLogin().UTC(), time.Millisecond, "LastLogin should be equal") - assert.WithinDurationf(t, peer.CreatedAt, storedPeer.CreatedAt.UTC(), time.Millisecond, "CreatedAt should be equal") - assert.Equal(t, peer.Ephemeral, storedPeer.Ephemeral) - assert.Equal(t, peer.Status.Connected, storedPeer.Status.Connected) - assert.Equal(t, peer.Status.LoginExpired, storedPeer.Status.LoginExpired) - assert.Equal(t, peer.Status.RequiresApproval, storedPeer.Status.RequiresApproval) - assert.WithinDurationf(t, peer.Status.LastSeen, storedPeer.Status.LastSeen.UTC(), time.Millisecond, "LastSeen should be equal") -} - -func TestSqlStore_GetPeerGroups(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - peerID := "cfefqs706sqkneg59g4g" - - groups, err := store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peerID) - require.NoError(t, err) - assert.Len(t, groups, 1) - assert.Equal(t, groups[0].Name, "All") - - err = store.AddPeerToGroup(context.Background(), accountID, peerID, "cfefqs706sqkneg59g4h") - require.NoError(t, err) - - groups, err = store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peerID) - require.NoError(t, err) - assert.Len(t, groups, 2) - - foreignPeerID := "foreign-peer" - err = store.AddPeerToGroup(context.Background(), accountID, foreignPeerID, "cfefqs706sqkneg59g4h") - require.NoError(t, err) - - groups, err = store.GetPeerGroups(context.Background(), LockingStrengthNone, "other-account", foreignPeerID) - require.NoError(t, err) - assert.Empty(t, groups, "groups of another account must not be returned") -} - -func TestSqlStore_GetAccountPeers(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - nameFilter string - ipFilter string - expectedCount int - }{ - { - name: "should retrieve peers for an existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectedCount: 5, - }, - { - name: "should return no peers for a non-existing account ID", - accountID: "nonexistent", - expectedCount: 0, - }, - { - name: "should return no peers for an empty account ID", - accountID: "", - expectedCount: 0, - }, - { - name: "should filter peers by name", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - nameFilter: "expiredhost", - expectedCount: 1, - }, - { - name: "should filter peers by partial name", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - nameFilter: "host", - expectedCount: 4, - }, - { - name: "should filter peers by ip", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - ipFilter: "100.64.39.54", - expectedCount: 1, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetAccountPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.nameFilter, tt.ipFilter) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - }) - } - -} - -func TestSqlStore_GetAccountPeersWithExpiration(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectedCount int - expectedPeerIDs []string - }{ - { - name: "should retrieve only non-expired peers with expiration enabled", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectedCount: 1, - expectedPeerIDs: []string{"notexpired01"}, - }, - { - name: "should return no peers with expiration for a non-existing account ID", - accountID: "nonexistent", - expectedCount: 0, - }, - { - name: "should return no peers with expiration for a empty account ID", - accountID: "", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetAccountPeersWithExpiration(context.Background(), LockingStrengthNone, tt.accountID) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - for i, peer := range peers { - assert.Equal(t, tt.expectedPeerIDs[i], peer.ID) - } - }) - } -} - -func TestSqlStore_GetAccountPeersWithExpiration_ExcludesAlreadyExpired(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - peers, err := store.GetAccountPeersWithExpiration(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - - // Verify the already-expired peer (cg05lnblo1hkg2j514p0) is not returned - for _, peer := range peers { - assert.NotEqual(t, "cg05lnblo1hkg2j514p0", peer.ID, "already expired peer should not be returned") - assert.False(t, peer.Status.LoginExpired, "returned peers should not have LoginExpired set") - } -} - -func TestSqlStore_GetAccountPeersWithInactivity(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectedCount int - }{ - { - name: "should retrieve peers with inactivity for an existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectedCount: 1, - }, - { - name: "should return no peers with inactivity for a non-existing account ID", - accountID: "nonexistent", - expectedCount: 0, - }, - { - name: "should return no peers with inactivity for an empty account ID", - accountID: "", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetAccountPeersWithInactivity(context.Background(), LockingStrengthNone, tt.accountID) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - }) - } -} - -func TestSqlStore_GetAllEphemeralPeers(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/storev1.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - peers, err := store.GetAllEphemeralPeers(context.Background(), LockingStrengthNone) - require.NoError(t, err) - require.Len(t, peers, 1) - require.True(t, peers[0].Ephemeral) -} - -func TestSqlStore_GetUserPeers(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - userID string - expectedCount int - }{ - { - name: "should retrieve peers for existing account ID and user ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - userID: "f4f6d672-63fb-11ec-90d6-0242ac120003", - expectedCount: 1, - }, - { - name: "should return no peers for non-existing account ID with existing user ID", - accountID: "nonexistent", - userID: "f4f6d672-63fb-11ec-90d6-0242ac120003", - expectedCount: 0, - }, - { - name: "should return no peers for non-existing user ID with existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - userID: "nonexistent_user", - expectedCount: 0, - }, - { - name: "should retrieve peers for another valid account ID and user ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - userID: "edafee4e-63fb-11ec-90d6-0242ac120003", - expectedCount: 3, - }, - { - name: "should return no peers for existing account ID with empty user ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - userID: "", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetUserPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.userID) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - }) - } -} - -func TestSqlStore_DeletePeer(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - peerID := "csrnkiq7qv9d8aitqd50" - - err = store.DeletePeer(context.Background(), accountID, peerID) - require.NoError(t, err) - - peer, err := store.GetPeerByID(context.Background(), LockingStrengthNone, accountID, peerID) - require.Error(t, err) - require.Nil(t, peer) -} - func TestSqlStore_DatabaseBlocking(t *testing.T) { store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) t.Cleanup(cleanup) @@ -3154,1074 +335,6 @@ func TestSqlStore_DatabaseBlocking(t *testing.T) { t.Logf("Test completed") } -func TestSqlStore_GetAccountCreatedBy(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectError bool - createdBy string - }{ - { - name: "existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectError: false, - createdBy: "edafee4e-63fb-11ec-90d6-0242ac120003", - }, - { - name: "non-existing account ID", - accountID: "nonexistent", - expectError: true, - }, - { - name: "empty account ID", - accountID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - createdBy, err := store.GetAccountCreatedBy(context.Background(), LockingStrengthNone, tt.accountID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Empty(t, createdBy) - } else { - require.NoError(t, err) - require.NotNil(t, createdBy) - require.Equal(t, tt.createdBy, createdBy) - } - }) - } - -} - -func TestSqlStore_GetUserByUserID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - userID string - expectError bool - }{ - { - name: "retrieve existing user", - userID: "edafee4e-63fb-11ec-90d6-0242ac120003", - expectError: false, - }, - { - name: "retrieve non-existing user", - userID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty user ID", - userID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - user, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, tt.userID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, user) - } else { - require.NoError(t, err) - require.NotNil(t, user) - require.Equal(t, tt.userID, user.Id) - } - }) - } -} - -func TestSqlStore_GetUserByPATID(t *testing.T) { - store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanUp) - assert.NoError(t, err) - - id := "9dj38s35-63fb-11ec-90d6-0242ac120003" - - user, err := store.GetUserByPATID(context.Background(), LockingStrengthNone, id) - require.NoError(t, err) - require.Equal(t, "f4f6d672-63fb-11ec-90d6-0242ac120003", user.Id) -} - -func TestSqlStore_SaveUser(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - user := &types.User{ - Id: "user-id", - AccountID: accountID, - Role: types.UserRoleAdmin, - IsServiceUser: false, - AutoGroups: []string{"groupA", "groupB"}, - Blocked: false, - LastLogin: util.ToPtr(time.Now().UTC()), - CreatedAt: time.Now().UTC().Add(-time.Hour), - Issued: types.UserIssuedIntegration, - } - err = store.SaveUser(context.Background(), user) - require.NoError(t, err) - - saveUser, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, user.Id) - require.NoError(t, err) - require.Equal(t, user.Id, saveUser.Id) - require.Equal(t, user.AccountID, saveUser.AccountID) - require.Equal(t, user.Role, saveUser.Role) - require.Equal(t, user.AutoGroups, saveUser.AutoGroups) - require.WithinDurationf(t, user.GetLastLogin(), saveUser.LastLogin.UTC(), time.Millisecond, "LastLogin should be equal") - require.WithinDurationf(t, user.CreatedAt, saveUser.CreatedAt.UTC(), time.Millisecond, "CreatedAt should be equal") - require.Equal(t, user.Issued, saveUser.Issued) - require.Equal(t, user.Blocked, saveUser.Blocked) - require.Equal(t, user.IsServiceUser, saveUser.IsServiceUser) -} - -func TestSqlStore_SaveUsers(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - accountUsers, err := store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Len(t, accountUsers, 2) - - users := []*types.User{ - { - Id: "user-1", - AccountID: accountID, - Issued: "api", - AutoGroups: []string{"groupA", "groupB"}, - }, - { - Id: "user-2", - AccountID: accountID, - Issued: "integration", - AutoGroups: []string{"groupA"}, - }, - } - err = store.SaveUsers(context.Background(), users) - require.NoError(t, err) - - accountUsers, err = store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Len(t, accountUsers, 4) - - users[1].AutoGroups = []string{"groupA", "groupC"} - err = store.SaveUsers(context.Background(), users) - require.NoError(t, err) - - user, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, users[1].Id) - require.NoError(t, err) - require.Equal(t, users[1].AutoGroups, user.AutoGroups) -} - -func TestSqlStore_SaveUserWithEncryption(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - // Enable encryption - key, err := crypt.GenerateKey() - require.NoError(t, err) - fieldEncrypt, err := crypt.NewFieldEncrypt(key) - require.NoError(t, err) - store.SetFieldEncrypt(fieldEncrypt) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - // rawUser is used to read raw (potentially encrypted) data from the database - // without any gorm hooks or automatic decryption - type rawUser struct { - Id string - Email string - Name string - } - - t.Run("save user with empty email and name", func(t *testing.T) { - user := &types.User{ - Id: "user-empty-fields", - AccountID: accountID, - Role: types.UserRoleUser, - Email: "", - Name: "", - AutoGroups: []string{"groupA"}, - } - err = store.SaveUser(context.Background(), user) - require.NoError(t, err) - - // Verify using direct database query that empty strings remain empty (not encrypted) - var raw rawUser - err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", user.Id).First(&raw).Error - require.NoError(t, err) - require.Equal(t, "", raw.Email, "empty email should remain empty in database") - require.Equal(t, "", raw.Name, "empty name should remain empty in database") - - // Verify manual decryption returns empty strings - decryptedEmail, err := fieldEncrypt.Decrypt(raw.Email) - require.NoError(t, err) - require.Equal(t, "", decryptedEmail) - - decryptedName, err := fieldEncrypt.Decrypt(raw.Name) - require.NoError(t, err) - require.Equal(t, "", decryptedName) - }) - - t.Run("save user with email and name", func(t *testing.T) { - user := &types.User{ - Id: "user-with-fields", - AccountID: accountID, - Role: types.UserRoleAdmin, - Email: "test@example.com", - Name: "Test User", - AutoGroups: []string{"groupB"}, - } - err = store.SaveUser(context.Background(), user) - require.NoError(t, err) - - // Verify using direct database query that the data is encrypted (not plaintext) - var raw rawUser - err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", user.Id).First(&raw).Error - require.NoError(t, err) - require.NotEqual(t, "test@example.com", raw.Email, "email should be encrypted in database") - require.NotEqual(t, "Test User", raw.Name, "name should be encrypted in database") - - // Verify manual decryption returns correct values - decryptedEmail, err := fieldEncrypt.Decrypt(raw.Email) - require.NoError(t, err) - require.Equal(t, "test@example.com", decryptedEmail) - - decryptedName, err := fieldEncrypt.Decrypt(raw.Name) - require.NoError(t, err) - require.Equal(t, "Test User", decryptedName) - }) - - t.Run("save multiple users with mixed fields", func(t *testing.T) { - users := []*types.User{ - { - Id: "batch-user-1", - AccountID: accountID, - Email: "", - Name: "", - }, - { - Id: "batch-user-2", - AccountID: accountID, - Email: "batch@example.com", - Name: "Batch User", - }, - } - err = store.SaveUsers(context.Background(), users) - require.NoError(t, err) - - // Verify first user (empty fields) using direct database query - var raw1 rawUser - err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", "batch-user-1").First(&raw1).Error - require.NoError(t, err) - require.Equal(t, "", raw1.Email, "empty email should remain empty in database") - require.Equal(t, "", raw1.Name, "empty name should remain empty in database") - - // Verify second user (with fields) using direct database query - var raw2 rawUser - err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", "batch-user-2").First(&raw2).Error - require.NoError(t, err) - require.NotEqual(t, "batch@example.com", raw2.Email, "email should be encrypted in database") - require.NotEqual(t, "Batch User", raw2.Name, "name should be encrypted in database") - - // Verify manual decryption returns empty strings for first user - decryptedEmail1, err := fieldEncrypt.Decrypt(raw1.Email) - require.NoError(t, err) - require.Equal(t, "", decryptedEmail1) - - decryptedName1, err := fieldEncrypt.Decrypt(raw1.Name) - require.NoError(t, err) - require.Equal(t, "", decryptedName1) - - // Verify manual decryption returns correct values for second user - decryptedEmail2, err := fieldEncrypt.Decrypt(raw2.Email) - require.NoError(t, err) - require.Equal(t, "batch@example.com", decryptedEmail2) - - decryptedName2, err := fieldEncrypt.Decrypt(raw2.Name) - require.NoError(t, err) - require.Equal(t, "Batch User", decryptedName2) - }) -} - -func TestSqlStore_DeleteUser(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" - - err = store.DeleteUser(context.Background(), accountID, userID) - require.NoError(t, err) - - user, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, userID) - require.Error(t, err) - require.Nil(t, user) - - userPATs, err := store.GetUserPATs(context.Background(), LockingStrengthNone, userID) - require.NoError(t, err) - require.Len(t, userPATs, 0) -} - -func TestSqlStore_GetPATByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" - - tests := []struct { - name string - patID string - expectError bool - }{ - { - name: "retrieve existing PAT", - patID: "9dj38s35-63fb-11ec-90d6-0242ac120003", - expectError: false, - }, - { - name: "retrieve non-existing PAT", - patID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty PAT ID", - patID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - pat, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, tt.patID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, pat) - } else { - require.NoError(t, err) - require.NotNil(t, pat) - require.Equal(t, tt.patID, pat.ID) - } - }) - } -} - -func TestSqlStore_GetUserPATs(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - userPATs, err := store.GetUserPATs(context.Background(), LockingStrengthNone, "f4f6d672-63fb-11ec-90d6-0242ac120003") - require.NoError(t, err) - require.Len(t, userPATs, 1) -} - -func TestSqlStore_GetPATByHashedToken(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - pat, err := store.GetPATByHashedToken(context.Background(), LockingStrengthNone, "SoMeHaShEdToKeN") - require.NoError(t, err) - require.Equal(t, "9dj38s35-63fb-11ec-90d6-0242ac120003", pat.ID) -} - -func TestSqlStore_MarkPATUsed(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" - patID := "9dj38s35-63fb-11ec-90d6-0242ac120003" - - err = store.MarkPATUsed(context.Background(), patID) - require.NoError(t, err) - - pat, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, patID) - require.NoError(t, err) - now := time.Now().UTC() - require.WithinRange(t, pat.LastUsed.UTC(), now.Add(-15*time.Second), now, "LastUsed should be within 1 second of now") -} - -func TestSqlStore_SavePAT(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - userID := "edafee4e-63fb-11ec-90d6-0242ac120003" - - pat := &types.PersonalAccessToken{ - ID: "pat-id", - UserID: userID, - Name: "token", - HashedToken: "SoMeHaShEdToKeN", - ExpirationDate: util.ToPtr(time.Now().UTC().Add(12 * time.Hour)), - CreatedBy: userID, - CreatedAt: time.Now().UTC().Add(time.Hour), - LastUsed: util.ToPtr(time.Now().UTC().Add(-15 * time.Minute)), - } - err = store.SavePAT(context.Background(), pat) - require.NoError(t, err) - - savePAT, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, pat.ID) - require.NoError(t, err) - require.Equal(t, pat.ID, savePAT.ID) - require.Equal(t, pat.UserID, savePAT.UserID) - require.Equal(t, pat.HashedToken, savePAT.HashedToken) - require.Equal(t, pat.CreatedBy, savePAT.CreatedBy) - require.WithinDurationf(t, pat.GetExpirationDate(), savePAT.ExpirationDate.UTC(), time.Millisecond, "ExpirationDate should be equal") - require.WithinDurationf(t, pat.CreatedAt, savePAT.CreatedAt.UTC(), time.Millisecond, "CreatedAt should be equal") - require.WithinDurationf(t, pat.GetLastUsed(), savePAT.LastUsed.UTC(), time.Millisecond, "LastUsed should be equal") -} - -func TestSqlStore_DeletePAT(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" - patID := "9dj38s35-63fb-11ec-90d6-0242ac120003" - - err = store.DeletePAT(context.Background(), userID, patID) - require.NoError(t, err) - - pat, err := store.GetPATByID(context.Background(), LockingStrengthNone, userID, patID) - require.Error(t, err) - require.Nil(t, pat) -} - -func TestSqlStore_SaveUsers_LargeBatch(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - accountUsers, err := store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Len(t, accountUsers, 2) - - usersToSave := make([]*types.User, 0) - - for i := 1; i <= 8000; i++ { - usersToSave = append(usersToSave, &types.User{ - Id: fmt.Sprintf("user-%d", i), - AccountID: accountID, - Role: types.UserRoleUser, - }) - } - - err = store.SaveUsers(context.Background(), usersToSave) - require.NoError(t, err) - - accountUsers, err = store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Equal(t, 8002, len(accountUsers)) -} - -func TestSqlStore_SaveGroups_LargeBatch(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - accountGroups, err := store.GetAccountGroups(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Len(t, accountGroups, 3) - - groupsToSave := make([]*types.Group, 0) - - for i := 1; i <= 8000; i++ { - groupsToSave = append(groupsToSave, &types.Group{ - ID: fmt.Sprintf("%d", i), - AccountID: accountID, - Name: fmt.Sprintf("group-%d", i), - }) - } - - err = store.CreateGroups(context.Background(), accountID, groupsToSave) - require.NoError(t, err) - - accountGroups, err = store.GetAccountGroups(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.Equal(t, 8003, len(accountGroups)) -} -func TestSqlStore_GetAccountRoutes(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - expectedCount int - }{ - { - name: "retrieve routes by existing account ID", - accountID: "bf1c8084-ba50-4ce7-9439-34653001fc3b", - expectedCount: 1, - }, - { - name: "non-existing account ID", - accountID: "nonexistent", - expectedCount: 0, - }, - { - name: "empty account ID", - accountID: "", - expectedCount: 0, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - routes, err := store.GetAccountRoutes(context.Background(), LockingStrengthNone, tt.accountID) - require.NoError(t, err) - require.Len(t, routes, tt.expectedCount) - }) - } -} - -func TestSqlStore_GetRouteByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - tests := []struct { - name string - routeID string - expectError bool - }{ - { - name: "retrieve existing route", - routeID: "ct03t427qv97vmtmglog", - expectError: false, - }, - { - name: "retrieve non-existing route", - routeID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty route ID", - routeID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - route, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, tt.routeID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, route) - } else { - require.NoError(t, err) - require.NotNil(t, route) - require.Equal(t, tt.routeID, string(route.ID)) - } - }) - } -} - -func TestSqlStore_SaveRoute(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - route := &nbroute.Route{ - ID: "route-id", - AccountID: accountID, - Network: netip.MustParsePrefix("10.10.0.0/16"), - NetID: "netID", - PeerGroups: []string{"routeA"}, - NetworkType: nbroute.IPv4Network, - Masquerade: true, - Metric: 9999, - Enabled: true, - Groups: []string{"groupA"}, - AccessControlGroups: []string{}, - } - err = store.SaveRoute(context.Background(), route) - require.NoError(t, err) - - saveRoute, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, string(route.ID)) - require.NoError(t, err) - require.Equal(t, route, saveRoute) - -} - -func TestSqlStore_DeleteRoute(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - routeID := "ct03t427qv97vmtmglog" - - err = store.DeleteRoute(context.Background(), accountID, routeID) - require.NoError(t, err) - - route, err := store.GetRouteByID(context.Background(), LockingStrengthNone, accountID, routeID) - require.Error(t, err) - require.Nil(t, route) -} - -func TestSqlStore_GetAccountMeta(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - accountMeta, err := store.GetAccountMeta(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.NotNil(t, accountMeta) - require.Equal(t, accountID, accountMeta.AccountID) - require.Equal(t, "edafee4e-63fb-11ec-90d6-0242ac120003", accountMeta.CreatedBy) - require.Equal(t, "test.com", accountMeta.Domain) - require.Equal(t, "private", accountMeta.DomainCategory) - require.Equal(t, time.Date(2024, time.October, 2, 14, 1, 38, 210000000, time.UTC), accountMeta.CreatedAt.UTC()) -} - -func TestSqlStore_GetAccountOnboarding(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "9439-34653001fc3b-bf1c8084-ba50-4ce7" - a, err := store.GetAccount(context.Background(), accountID) - require.NoError(t, err) - t.Logf("Onboarding: %+v", a.Onboarding) - err = store.SaveAccount(context.Background(), a) - require.NoError(t, err) - onboarding, err := store.GetAccountOnboarding(context.Background(), accountID) - require.NoError(t, err) - require.NotNil(t, onboarding) - require.Equal(t, accountID, onboarding.AccountID) - require.Equal(t, time.Date(2024, time.October, 2, 14, 1, 38, 210000000, time.UTC), onboarding.CreatedAt.UTC()) -} - -func TestSqlStore_SaveAccountOnboarding(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - t.Run("New onboarding should be saved correctly", func(t *testing.T) { - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - onboarding := &types.AccountOnboarding{ - AccountID: accountID, - SignupFormPending: true, - OnboardingFlowPending: true, - } - - err = store.SaveAccountOnboarding(context.Background(), onboarding) - require.NoError(t, err) - - savedOnboarding, err := store.GetAccountOnboarding(context.Background(), accountID) - require.NoError(t, err) - require.Equal(t, onboarding.SignupFormPending, savedOnboarding.SignupFormPending) - require.Equal(t, onboarding.OnboardingFlowPending, savedOnboarding.OnboardingFlowPending) - }) - - t.Run("Existing onboarding should be updated correctly", func(t *testing.T) { - accountID := "9439-34653001fc3b-bf1c8084-ba50-4ce7" - onboarding, err := store.GetAccountOnboarding(context.Background(), accountID) - require.NoError(t, err) - - onboarding.OnboardingFlowPending = !onboarding.OnboardingFlowPending - onboarding.SignupFormPending = !onboarding.SignupFormPending - - err = store.SaveAccountOnboarding(context.Background(), onboarding) - require.NoError(t, err) - - savedOnboarding, err := store.GetAccountOnboarding(context.Background(), accountID) - require.NoError(t, err) - require.Equal(t, onboarding.SignupFormPending, savedOnboarding.SignupFormPending) - require.Equal(t, onboarding.OnboardingFlowPending, savedOnboarding.OnboardingFlowPending) - }) -} - -func TestSqlStore_GetAnyAccountID(t *testing.T) { - t.Run("should return account ID when accounts exist", func(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID, err := store.GetAnyAccountID(context.Background()) - require.NoError(t, err) - assert.Equal(t, "bf1c8084-ba50-4ce7-9439-34653001fc3b", accountID) - }) - - t.Run("should return error when no accounts exist", func(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID, err := store.GetAnyAccountID(context.Background()) - require.Error(t, err) - sErr, ok := status.FromError(err) - assert.True(t, ok) - assert.Equal(t, sErr.Type(), status.NotFound) - assert.Empty(t, accountID) - }) -} - -func BenchmarkGetAccountPeers(b *testing.B) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", b.TempDir()) - if err != nil { - b.Fatal(err) - } - b.Cleanup(cleanup) - - numberOfPeers := 1000 - numberOfGroups := 200 - numberOfPeersPerGroup := 500 - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - peers := make([]*nbpeer.Peer, 0, numberOfPeers) - for i := 0; i < numberOfPeers; i++ { - peer := &nbpeer.Peer{ - ID: fmt.Sprintf("peer-%d", i), - AccountID: accountID, - Key: fmt.Sprintf("key-%d", i), - DNSLabel: fmt.Sprintf("peer%d.example.com", i), - IP: intToIPv4(uint32(i)), - } - err = store.AddPeerToAccount(context.Background(), peer) - if err != nil { - b.Fatalf("Failed to add peer: %v", err) - } - peers = append(peers, peer) - } - - for i := 0; i < numberOfGroups; i++ { - groupID := fmt.Sprintf("group-%d", i) - group := &types.Group{ - ID: groupID, - AccountID: accountID, - } - err = store.CreateGroup(context.Background(), group) - if err != nil { - b.Fatalf("Failed to create group: %v", err) - } - for j := 0; j < numberOfPeersPerGroup; j++ { - peerIndex := (i*numberOfPeersPerGroup + j) % numberOfPeers - err = store.AddPeerToGroup(context.Background(), accountID, peers[peerIndex].ID, groupID) - if err != nil { - b.Fatalf("Failed to add peer to group: %v", err) - } - } - } - - b.ResetTimer() - for i := 0; i < b.N; i++ { - _, err := store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peers[i%numberOfPeers].ID) - if err != nil { - b.Fatal(err) - } - } -} - -func intToIPv4(n uint32) netip.Addr { - var b [4]byte - binary.BigEndian.PutUint32(b[:], n) - return netip.AddrFrom4(b) -} - -func TestSqlStore_GetPeersByGroupIDs(t *testing.T) { - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - group1ID := "test-group-1" - group2ID := "test-group-2" - emptyGroupID := "empty-group" - - peer1 := "cfefqs706sqkneg59g4g" - peer2 := "cfeg6sf06sqkneg59g50" - - tests := []struct { - name string - groupIDs []string - expectedPeers []string - expectedCount int - }{ - { - name: "retrieve peers from single group with multiple peers", - groupIDs: []string{group1ID}, - expectedPeers: []string{peer1, peer2}, - expectedCount: 2, - }, - { - name: "retrieve peers from single group with one peer", - groupIDs: []string{group2ID}, - expectedPeers: []string{peer1}, - expectedCount: 1, - }, - { - name: "retrieve peers from multiple groups (with overlap)", - groupIDs: []string{group1ID, group2ID}, - expectedPeers: []string{peer1, peer2}, // should deduplicate - expectedCount: 2, - }, - { - name: "retrieve peers from existing 'All' group", - groupIDs: []string{"cfefqs706sqkneg59g3g"}, // All group from test data - expectedPeers: []string{peer1, peer2}, - expectedCount: 2, - }, - { - name: "retrieve peers from empty group", - groupIDs: []string{emptyGroupID}, - expectedPeers: []string{}, - expectedCount: 0, - }, - { - name: "retrieve peers from non-existing group", - groupIDs: []string{"non-existing-group"}, - expectedPeers: []string{}, - expectedCount: 0, - }, - { - name: "empty group IDs list", - groupIDs: []string{}, - expectedPeers: []string{}, - expectedCount: 0, - }, - { - name: "mix of existing and non-existing groups", - groupIDs: []string{group1ID, "non-existing-group"}, - expectedPeers: []string{peer1, peer2}, - expectedCount: 2, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_policy_migrate.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - ctx := context.Background() - - groups := []*types.Group{ - { - ID: group1ID, - AccountID: accountID, - }, - { - ID: group2ID, - AccountID: accountID, - }, - } - require.NoError(t, store.CreateGroups(ctx, accountID, groups)) - - otherAccount := newAccountWithId(ctx, "other-account", "other-user", "") - require.NoError(t, store.SaveAccount(ctx, otherAccount)) - foreignPeer := &nbpeer.Peer{ID: "foreign-peer", AccountID: otherAccount.Id} - require.NoError(t, store.AddPeerToAccount(ctx, foreignPeer)) - - require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer1, group1ID)) - require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer2, group1ID)) - require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer1, group2ID)) - require.NoError(t, store.AddPeerToGroup(ctx, accountID, foreignPeer.ID, group1ID)) - - peers, err := store.GetPeersByGroupIDs(ctx, accountID, tt.groupIDs) - require.NoError(t, err) - require.Len(t, peers, tt.expectedCount) - - if tt.expectedCount > 0 { - actualPeerIDs := make([]string, len(peers)) - for i, peer := range peers { - actualPeerIDs[i] = peer.ID - } - assert.ElementsMatch(t, tt.expectedPeers, actualPeerIDs) - - // Verify all returned peers belong to the correct account - for _, peer := range peers { - assert.Equal(t, accountID, peer.AccountID) - } - } - }) - } -} - -func TestSqlStore_GetUserIDByPeerKey(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - userID := "test-user-123" - peerKey := "peer-key-abc" - - peer := &nbpeer.Peer{ - ID: "test-peer-1", - Key: peerKey, - AccountID: existingAccountID, - UserID: userID, - IP: netip.AddrFrom4([4]byte{10, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::a00:1"), - DNSLabel: "test-peer-1", - } - - err = store.AddPeerToAccount(context.Background(), peer) - require.NoError(t, err) - - retrievedUserID, err := store.GetUserIDByPeerKey(context.Background(), LockingStrengthNone, peerKey) - require.NoError(t, err) - assert.Equal(t, userID, retrievedUserID) -} - -func TestSqlStore_GetUserIDByPeerKey_NotFound(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - nonExistentPeerKey := "non-existent-peer-key" - - userID, err := store.GetUserIDByPeerKey(context.Background(), LockingStrengthNone, nonExistentPeerKey) - require.Error(t, err) - assert.Equal(t, "", userID) -} - -func TestSqlStore_GetUserIDByPeerKey_NoUserID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - existingAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - peerKey := "peer-key-abc" - - peer := &nbpeer.Peer{ - ID: "test-peer-1", - Key: peerKey, - AccountID: existingAccountID, - UserID: "", - IP: netip.AddrFrom4([4]byte{10, 0, 0, 1}), - IPv6: netip.MustParseAddr("fd00::a00:1"), - DNSLabel: "test-peer-1", - } - - err = store.AddPeerToAccount(context.Background(), peer) - require.NoError(t, err) - - retrievedUserID, err := store.GetUserIDByPeerKey(context.Background(), LockingStrengthNone, peerKey) - require.NoError(t, err) - assert.Equal(t, "", retrievedUserID) -} - -func TestSqlStore_ApproveAccountPeers(t *testing.T) { - runTestForAllEngines(t, "", func(t *testing.T, store Store) { - accountID := "test-account" - ctx := context.Background() - - account := newAccountWithId(ctx, accountID, "testuser", "example.com") - err := store.SaveAccount(ctx, account) - require.NoError(t, err) - - peers := []*nbpeer.Peer{ - { - ID: "peer1", - AccountID: accountID, - DNSLabel: "peer1.netbird.cloud", - Key: "peer1-key", - IP: netip.MustParseAddr("100.64.0.1"), - IPv6: netip.MustParseAddr("fd00::1"), - Status: &nbpeer.PeerStatus{ - RequiresApproval: true, - LastSeen: time.Now().UTC(), - }, - }, - { - ID: "peer2", - AccountID: accountID, - DNSLabel: "peer2.netbird.cloud", - Key: "peer2-key", - IP: netip.MustParseAddr("100.64.0.2"), - IPv6: netip.MustParseAddr("fd00::2"), - Status: &nbpeer.PeerStatus{ - RequiresApproval: true, - LastSeen: time.Now().UTC(), - }, - }, - { - ID: "peer3", - AccountID: accountID, - DNSLabel: "peer3.netbird.cloud", - Key: "peer3-key", - IP: netip.MustParseAddr("100.64.0.3"), - IPv6: netip.MustParseAddr("fd00::3"), - Status: &nbpeer.PeerStatus{ - RequiresApproval: false, - LastSeen: time.Now().UTC(), - }, - }, - } - - for _, peer := range peers { - err = store.AddPeerToAccount(ctx, peer) - require.NoError(t, err) - } - - t.Run("approve all pending peers", func(t *testing.T) { - count, err := store.ApproveAccountPeers(ctx, accountID) - require.NoError(t, err) - assert.Equal(t, 2, count) - - allPeers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "") - require.NoError(t, err) - - for _, peer := range allPeers { - assert.False(t, peer.Status.RequiresApproval, "peer %s should not require approval", peer.ID) - } - }) - - t.Run("no peers to approve", func(t *testing.T) { - count, err := store.ApproveAccountPeers(ctx, accountID) - require.NoError(t, err) - assert.Equal(t, 0, count) - }) - - t.Run("non-existent account", func(t *testing.T) { - count, err := store.ApproveAccountPeers(ctx, "non-existent") - require.NoError(t, err) - assert.Equal(t, 0, count) - }) - }) -} - func TestSqlStore_ExecuteInTransaction_Timeout(t *testing.T) { if os.Getenv("NETBIRD_STORE_ENGINE") == "mysql" { t.Skip("Skipping timeout test for MySQL") @@ -4235,7 +348,6 @@ func TestSqlStore_ExecuteInTransaction_Timeout(t *testing.T) { sqlStore, ok := store.(*SqlStore) require.True(t, ok) - assert.Equal(t, 1*time.Second, sqlStore.transactionTimeout) ctx := context.Background() err = sqlStore.ExecuteInTransaction(ctx, func(transaction Store) error { @@ -4249,479 +361,6 @@ func TestSqlStore_ExecuteInTransaction_Timeout(t *testing.T) { assert.Contains(t, err.Error(), "transaction has already been committed or rolled back", "expected transaction rolled back error, got: %v", err) } -func TestSqlStore_CreateZone(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - savedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, accountID, zone.ID) - require.NoError(t, err) - require.NotNil(t, savedZone) - assert.Equal(t, zone.ID, savedZone.ID) - assert.Equal(t, zone.Name, savedZone.Name) - assert.Equal(t, zone.Domain, savedZone.Domain) - assert.Equal(t, zone.Enabled, savedZone.Enabled) - assert.Equal(t, zone.EnableSearchDomain, savedZone.EnableSearchDomain) - assert.Equal(t, zone.DistributionGroups, savedZone.DistributionGroups) -} - -func TestSqlStore_GetZoneByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - zoneID string - expectError bool - }{ - { - name: "retrieve existing zone", - accountID: accountID, - zoneID: zone.ID, - expectError: false, - }, - { - name: "retrieve non-existing zone", - accountID: accountID, - zoneID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty zone ID", - accountID: accountID, - zoneID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - savedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, tt.accountID, tt.zoneID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, savedZone) - } else { - require.NoError(t, err) - require.NotNil(t, savedZone) - assert.Equal(t, tt.zoneID, savedZone.ID) - } - }) - } -} - -func TestSqlStore_GetAccountZones(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone1 := zones.NewZone(accountID, "Zone 1", "example1.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone1) - require.NoError(t, err) - - zone2 := zones.NewZone(accountID, "Zone 2", "example2.com", true, true, []string{"group1", "group2"}) - err = store.CreateZone(context.Background(), zone2) - require.NoError(t, err) - - allZones, err := store.GetAccountZones(context.Background(), LockingStrengthNone, accountID) - require.NoError(t, err) - require.NotNil(t, allZones) - assert.GreaterOrEqual(t, len(allZones), 2) - - zoneIDs := make(map[string]bool) - for _, z := range allZones { - zoneIDs[z.ID] = true - } - assert.True(t, zoneIDs[zone1.ID]) - assert.True(t, zoneIDs[zone2.ID]) -} - -func TestSqlStore_GetZoneByDomain(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - otherAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3c" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - domain string - expectError bool - errorType status.Type - }{ - { - name: "retrieve existing zone by domain", - accountID: accountID, - domain: "example.com", - expectError: false, - }, - { - name: "retrieve non-existing zone domain", - accountID: accountID, - domain: "non-existing.com", - expectError: true, - errorType: status.NotFound, - }, - { - name: "retrieve with empty domain", - accountID: accountID, - domain: "", - expectError: true, - errorType: status.NotFound, - }, - { - name: "retrieve with different account ID", - accountID: otherAccountID, - domain: "example.com", - expectError: true, - errorType: status.NotFound, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - savedZone, err := store.GetZoneByDomain(context.Background(), tt.accountID, tt.domain) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, tt.errorType, sErr.Type()) - require.Nil(t, savedZone) - } else { - require.NoError(t, err) - require.NotNil(t, savedZone) - assert.Equal(t, tt.domain, savedZone.Domain) - assert.Equal(t, zone.ID, savedZone.ID) - assert.Equal(t, zone.Name, savedZone.Name) - } - }) - } -} - -func TestSqlStore_UpdateZone(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - zone.Name = "Updated Zone" - zone.Domain = "updated.com" - zone.Enabled = false - zone.EnableSearchDomain = true - zone.DistributionGroups = []string{"group2", "group3"} - - err = store.UpdateZone(context.Background(), zone) - require.NoError(t, err) - - updatedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, accountID, zone.ID) - require.NoError(t, err) - require.NotNil(t, updatedZone) - assert.Equal(t, "Updated Zone", updatedZone.Name) - assert.Equal(t, "updated.com", updatedZone.Domain) - assert.False(t, updatedZone.Enabled) - assert.True(t, updatedZone.EnableSearchDomain) - assert.Equal(t, []string{"group2", "group3"}, updatedZone.DistributionGroups) -} - -func TestSqlStore_DeleteZone(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - err = store.DeleteZone(context.Background(), accountID, zone.ID) - require.NoError(t, err) - - deletedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, accountID, zone.ID) - require.Error(t, err) - require.Nil(t, deletedZone) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) -} - -func TestSqlStore_CreateDNSRecord(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - - err = store.CreateDNSRecord(context.Background(), record) - require.NoError(t, err) - - savedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, accountID, zone.ID, record.ID) - require.NoError(t, err) - require.NotNil(t, savedRecord) - assert.Equal(t, record.ID, savedRecord.ID) - assert.Equal(t, record.Name, savedRecord.Name) - assert.Equal(t, record.Type, savedRecord.Type) - assert.Equal(t, record.Content, savedRecord.Content) - assert.Equal(t, record.TTL, savedRecord.TTL) - assert.Equal(t, zone.ID, savedRecord.ZoneID) -} - -func TestSqlStore_GetDNSRecordByID(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - err = store.CreateDNSRecord(context.Background(), record) - require.NoError(t, err) - - tests := []struct { - name string - accountID string - zoneID string - recordID string - expectError bool - }{ - { - name: "retrieve existing record", - accountID: accountID, - zoneID: zone.ID, - recordID: record.ID, - expectError: false, - }, - { - name: "retrieve non-existing record", - accountID: accountID, - zoneID: zone.ID, - recordID: "non-existing", - expectError: true, - }, - { - name: "retrieve with empty record ID", - accountID: accountID, - zoneID: zone.ID, - recordID: "", - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - savedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, tt.accountID, tt.zoneID, tt.recordID) - if tt.expectError { - require.Error(t, err) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) - require.Nil(t, savedRecord) - } else { - require.NoError(t, err) - require.NotNil(t, savedRecord) - assert.Equal(t, tt.recordID, savedRecord.ID) - } - }) - } -} - -func TestSqlStore_GetZoneDNSRecords(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - recordA := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - err = store.CreateDNSRecord(context.Background(), recordA) - require.NoError(t, err) - - recordAAAA := records.NewRecord(accountID, zone.ID, "ipv6.example.com", records.RecordTypeAAAA, "2001:db8::1", 300) - err = store.CreateDNSRecord(context.Background(), recordAAAA) - require.NoError(t, err) - - recordCNAME := records.NewRecord(accountID, zone.ID, "alias.example.com", records.RecordTypeCNAME, "www.example.com", 300) - err = store.CreateDNSRecord(context.Background(), recordCNAME) - require.NoError(t, err) - - allRecords, err := store.GetZoneDNSRecords(context.Background(), LockingStrengthNone, accountID, zone.ID) - require.NoError(t, err) - require.NotNil(t, allRecords) - assert.Equal(t, 3, len(allRecords)) - - recordIDs := make(map[string]bool) - for _, r := range allRecords { - recordIDs[r.ID] = true - } - assert.True(t, recordIDs[recordA.ID]) - assert.True(t, recordIDs[recordAAAA.ID]) - assert.True(t, recordIDs[recordCNAME.ID]) -} - -func TestSqlStore_GetZoneDNSRecordsByName(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - record1 := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - err = store.CreateDNSRecord(context.Background(), record1) - require.NoError(t, err) - - record2 := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeAAAA, "2001:db8::1", 300) - err = store.CreateDNSRecord(context.Background(), record2) - require.NoError(t, err) - - record3 := records.NewRecord(accountID, zone.ID, "mail.example.com", records.RecordTypeA, "192.168.1.2", 600) - err = store.CreateDNSRecord(context.Background(), record3) - require.NoError(t, err) - - recordsByName, err := store.GetZoneDNSRecordsByName(context.Background(), LockingStrengthNone, accountID, zone.ID, "www.example.com") - require.NoError(t, err) - require.NotNil(t, recordsByName) - assert.Equal(t, 2, len(recordsByName)) - - for _, r := range recordsByName { - assert.Equal(t, "www.example.com", r.Name) - } -} - -func TestSqlStore_UpdateDNSRecord(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - err = store.CreateDNSRecord(context.Background(), record) - require.NoError(t, err) - - record.Name = "api.example.com" - record.Content = "192.168.1.100" - record.TTL = 600 - - err = store.UpdateDNSRecord(context.Background(), record) - require.NoError(t, err) - - updatedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, accountID, zone.ID, record.ID) - require.NoError(t, err) - require.NotNil(t, updatedRecord) - assert.Equal(t, "api.example.com", updatedRecord.Name) - assert.Equal(t, "192.168.1.100", updatedRecord.Content) - assert.Equal(t, 600, updatedRecord.TTL) -} - -func TestSqlStore_DeleteDNSRecord(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - record := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - err = store.CreateDNSRecord(context.Background(), record) - require.NoError(t, err) - - err = store.DeleteDNSRecord(context.Background(), accountID, zone.ID, record.ID) - require.NoError(t, err) - - deletedRecord, err := store.GetDNSRecordByID(context.Background(), LockingStrengthNone, accountID, zone.ID, record.ID) - require.Error(t, err) - require.Nil(t, deletedRecord) - sErr, ok := status.FromError(err) - require.True(t, ok) - require.Equal(t, sErr.Type(), status.NotFound) -} - -func TestSqlStore_DeleteZoneDNSRecords(t *testing.T) { - store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) - t.Cleanup(cleanup) - require.NoError(t, err) - - accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" - - zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) - err = store.CreateZone(context.Background(), zone) - require.NoError(t, err) - - record1 := records.NewRecord(accountID, zone.ID, "www.example.com", records.RecordTypeA, "192.168.1.1", 300) - err = store.CreateDNSRecord(context.Background(), record1) - require.NoError(t, err) - - record2 := records.NewRecord(accountID, zone.ID, "mail.example.com", records.RecordTypeA, "192.168.1.2", 600) - err = store.CreateDNSRecord(context.Background(), record2) - require.NoError(t, err) - - allRecords, err := store.GetZoneDNSRecords(context.Background(), LockingStrengthNone, accountID, zone.ID) - require.NoError(t, err) - assert.Equal(t, 2, len(allRecords)) - - err = store.DeleteZoneDNSRecords(context.Background(), accountID, zone.ID) - require.NoError(t, err) - - remainingRecords, err := store.GetZoneDNSRecords(context.Background(), LockingStrengthNone, accountID, zone.ID) - require.NoError(t, err) - assert.Equal(t, 0, len(remainingRecords)) -} - // TestNewSqliteStore_BusyTimeoutApplied opens a fresh SQLite store and verifies // that the _busy_timeout DSN parameter took effect at the driver level. Without // this, lock contention on the single SQLite connection waits indefinitely on @@ -4773,3 +412,88 @@ func TestNewSqliteStore_BusyTimeoutRespectsUserOverride(t *testing.T) { }) } } + +func TestSqlStore_ExecuteInTransaction_RestoresForeignKeyChecksOnMysql(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + sqlStore := store.(*SqlStore) + if sqlStore.conn.Engine() != types.MysqlStoreEngine { + t.Skip("FOREIGN_KEY_CHECKS is MySQL specific") + } + sqlDB, err := sqlStore.GetDB().DB() + require.NoError(t, err) + sqlDB.SetMaxOpenConns(1) + ctx := context.Background() + + foreignKeyChecks := func() int { + var enabled int + require.NoError(t, sqlStore.GetDB().Raw("SELECT @@SESSION.foreign_key_checks").Scan(&enabled).Error) + return enabled + } + + err = store.ExecuteInTransaction(ctx, func(Store) error { return assert.AnError }) + require.ErrorIs(t, err, assert.AnError) + assert.Equal(t, 1, foreignKeyChecks()) + + require.Panics(t, func() { + _ = store.ExecuteInTransaction(ctx, func(Store) error { panic("boom") }) + }) + assert.Equal(t, 1, foreignKeyChecks()) + + err = sqlStore.transaction(ctx, func(*gorm.DB) error { return assert.AnError }) + require.ErrorIs(t, err, assert.AnError) + assert.Equal(t, 1, foreignKeyChecks()) + + err = store.ExecuteInTransaction(ctx, func(transaction Store) error { + bound := transaction.(*SqlStore) + require.NoError(t, bound.transaction(ctx, func(*gorm.DB) error { return nil })) + var enabled int + require.NoError(t, bound.GetDB().Raw("SELECT @@SESSION.foreign_key_checks").Scan(&enabled).Error) + assert.Equal(t, 0, enabled, "a savepoint must not re-enable FK checks for the rest of the transaction") + return nil + }) + require.NoError(t, err) + assert.Equal(t, 1, foreignKeyChecks()) + }) +} + +func TestSqlStore_Transaction_RollsBackOnError(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + ctx := context.Background() + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + group := &types.Group{ID: "rolled-back-group", AccountID: accountID, Name: "rolled back", Issued: "api"} + + err := store.(*SqlStore).transaction(ctx, func(tx *gorm.DB) error { + require.NoError(t, tx.Omit(clause.Associations).Create(group).Error) + return assert.AnError + }) + require.ErrorIs(t, err, assert.AnError) + + _, err = store.GetGroupByID(ctx, LockingStrengthNone, accountID, group.ID) + require.Error(t, err) + }) +} + +func TestSqlStore_Transaction_NestedIsSavepoint(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + ctx := context.Background() + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + outer := &types.Group{ID: "outer-group", AccountID: accountID, Name: "outer", Issued: "api"} + inner := &types.Group{ID: "inner-group", AccountID: accountID, Name: "inner", Issued: "api"} + + err := store.ExecuteInTransaction(ctx, func(transaction Store) error { + require.NoError(t, transaction.CreateGroup(ctx, outer)) + err := transaction.(*SqlStore).transaction(ctx, func(tx *gorm.DB) error { + require.NoError(t, tx.Omit(clause.Associations).Create(inner).Error) + return assert.AnError + }) + require.ErrorIs(t, err, assert.AnError) + return nil + }) + require.NoError(t, err) + + _, err = store.GetGroupByID(ctx, LockingStrengthNone, accountID, outer.ID) + require.NoError(t, err) + _, err = store.GetGroupByID(ctx, LockingStrengthNone, accountID, inner.ID) + require.Error(t, err) + }) +} diff --git a/management/server/store/sql_store_user.go b/management/server/store/sql_store_user.go new file mode 100644 index 000000000..4db03fb71 --- /dev/null +++ b/management/server/store/sql_store_user.go @@ -0,0 +1,272 @@ +package store + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "time" + + "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// SaveUsers saves the given list of users to the database. +func (s *SqlStore) SaveUsers(ctx context.Context, users []*types.User) error { + if len(users) == 0 { + return nil + } + + usersCopy := make([]*types.User, len(users)) + for i, user := range users { + userCopy := user.Copy() + userCopy.Email = user.Email + userCopy.Name = user.Name + if err := userCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { + return fmt.Errorf("encrypt user: %w", err) + } + usersCopy[i] = userCopy + } + + result := s.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&usersCopy) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save users to store: %s", result.Error) + return status.Errorf(status.Internal, "failed to save users to store") + } + return nil +} + +// SaveUser saves the given user to the database. +func (s *SqlStore) SaveUser(ctx context.Context, user *types.User) error { + userCopy := user.Copy() + userCopy.Email = user.Email + userCopy.Name = user.Name + + if err := userCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { + return fmt.Errorf("encrypt user: %w", err) + } + + result := s.db.Save(userCopy) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save user to store: %s", result.Error) + return status.Errorf(status.Internal, "failed to save user to store") + } + return nil +} + +func (s *SqlStore) GetUserByPATID(ctx context.Context, lockStrength LockingStrength, patID string) (*types.User, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var user types.User + result := tx. + Joins("JOIN personal_access_tokens ON personal_access_tokens.user_id = users.id"). + Where("personal_access_tokens.id = ?", patID).Take(&user) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewPATNotFoundError(patID) + } + log.WithContext(ctx).Errorf("failed to get token user from the store: %s", result.Error) + return nil, status.NewGetUserFromStoreError() + } + + if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt user: %w", err) + } + + return &user, nil +} + +func (s *SqlStore) GetUserByUserID(ctx context.Context, lockStrength LockingStrength, userID string) (*types.User, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var user types.User + result := tx.Take(&user, idQueryCondition, userID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewUserNotFoundError(userID) + } + return nil, status.NewGetUserFromStoreError() + } + + if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt user: %w", err) + } + + return &user, nil +} + +func (s *SqlStore) DeleteUser(ctx context.Context, accountID, userID string) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { + result := tx.Delete(&types.PersonalAccessToken{}, "user_id = ?", userID) + if result.Error != nil { + return result.Error + } + + return tx.Delete(&types.User{}, accountAndIDQueryCondition, accountID, userID).Error + }) + if err != nil { + log.WithContext(ctx).Errorf("failed to delete user from the store: %s", err) + return status.Errorf(status.Internal, "failed to delete user from store") + } + + return nil +} + +func (s *SqlStore) GetAccountUsers(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.User, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var users []*types.User + result := tx.Find(&users, accountIDCondition, accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "accountID not found: index lookup failed") + } + log.WithContext(ctx).Errorf("error when getting users from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "issue getting users from store") + } + + for _, user := range users { + if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt user: %w", err) + } + } + + return users, nil +} + +func (s *SqlStore) GetAccountOwner(ctx context.Context, lockStrength LockingStrength, accountID string) (*types.User, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var user types.User + result := tx.Take(&user, "account_id = ? AND role = ?", accountID, types.UserRoleOwner) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "account owner not found: index lookup failed") + } + return nil, status.Errorf(status.Internal, "failed to get account owner from the store") + } + + if err := user.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt user: %w", err) + } + + return &user, nil +} + +func (s *SqlStore) getUsers(ctx context.Context, accountID string) ([]types.User, error) { + const query = `SELECT id, account_id, role, is_service_user, non_deletable, service_user_name, auto_groups, blocked, pending_approval, last_login, created_at, issued, integration_ref_id, integration_ref_integration_type, email, name FROM users WHERE account_id = $1` + rows, err := s.pgxPool().Query(ctx, query, accountID) + if err != nil { + return nil, err + } + users, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (types.User, error) { + var u types.User + var autoGroups []byte + var lastLogin, createdAt sql.NullTime + var isServiceUser, nonDeletable, blocked, pendingApproval sql.NullBool + err := row.Scan(&u.Id, &u.AccountID, &u.Role, &isServiceUser, &nonDeletable, &u.ServiceUserName, &autoGroups, &blocked, &pendingApproval, &lastLogin, &createdAt, &u.Issued, &u.IntegrationReference.ID, &u.IntegrationReference.IntegrationType, &u.Email, &u.Name) + if err == nil { + if lastLogin.Valid { + u.LastLogin = &lastLogin.Time + } + if createdAt.Valid { + u.CreatedAt = createdAt.Time + } + if isServiceUser.Valid { + u.IsServiceUser = isServiceUser.Bool + } + if nonDeletable.Valid { + u.NonDeletable = nonDeletable.Bool + } + if blocked.Valid { + u.Blocked = blocked.Bool + } + if pendingApproval.Valid { + u.PendingApproval = pendingApproval.Bool + } + if autoGroups != nil { + _ = json.Unmarshal(autoGroups, &u.AutoGroups) + } else { + u.AutoGroups = []string{} + } + } + return u, err + }) + if err != nil { + return nil, err + } + return users, nil +} + +func (s *SqlStore) GetAccountByUser(ctx context.Context, userID string) (*types.Account, error) { + var user types.User + result := s.db.Select("account_id").Take(&user, idQueryCondition, userID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + return nil, status.NewGetAccountFromStoreError(result.Error) + } + + if user.AccountID == "" { + return nil, status.Errorf(status.NotFound, "account not found: index lookup failed") + } + + return s.GetAccount(ctx, user.AccountID) +} + +func (s *SqlStore) GetAccountIDByUserID(ctx context.Context, lockStrength LockingStrength, userID string) (string, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var accountID string + result := tx.Model(&types.User{}). + Select("account_id").Where(idQueryCondition, userID).Take(&accountID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return "", status.Errorf(status.NotFound, "account not found: index lookup failed") + } + return "", status.NewGetAccountFromStoreError(result.Error) + } + + return accountID, nil +} + +// SaveUserLastLogin stores the last login time for a user in DB. +func (s *SqlStore) SaveUserLastLogin(ctx context.Context, accountID, userID string, lastLogin time.Time) error { + var user types.User + result := s.db.Take(&user, accountAndIDQueryCondition, accountID, userID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return status.NewUserNotFoundError(userID) + } + return status.NewGetUserFromStoreError() + } + + if !lastLogin.IsZero() { + user.LastLogin = &lastLogin + return s.db.Save(&user).Error + } + + return nil +} diff --git a/management/server/store/sql_store_user_invite.go b/management/server/store/sql_store_user_invite.go new file mode 100644 index 000000000..3c93a0d19 --- /dev/null +++ b/management/server/store/sql_store_user_invite.go @@ -0,0 +1,139 @@ +package store + +import ( + "context" + "errors" + "fmt" + "strings" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// SaveUserInvite saves a user invite to the database +func (s *SqlStore) SaveUserInvite(ctx context.Context, invite *types.UserInviteRecord) error { + inviteCopy := invite.Copy() + if err := inviteCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil { + return fmt.Errorf("encrypt invite: %w", err) + } + + result := s.db.Save(inviteCopy) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to save user invite to store: %s", result.Error) + return status.Errorf(status.Internal, "failed to save user invite to store") + } + return nil +} + +// GetUserInviteByID retrieves a user invite by its ID and account ID +func (s *SqlStore) GetUserInviteByID(ctx context.Context, lockStrength LockingStrength, accountID, inviteID string) (*types.UserInviteRecord, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var invite types.UserInviteRecord + result := tx.Where("account_id = ?", accountID).Take(&invite, idQueryCondition, inviteID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "user invite not found") + } + log.WithContext(ctx).Errorf("failed to get user invite from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get user invite from store") + } + + if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt invite: %w", err) + } + + return &invite, nil +} + +// GetUserInviteByHashedToken retrieves a user invite by its hashed token +func (s *SqlStore) GetUserInviteByHashedToken(ctx context.Context, lockStrength LockingStrength, hashedToken string) (*types.UserInviteRecord, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var invite types.UserInviteRecord + result := tx.Take(&invite, "hashed_token = ?", hashedToken) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "user invite not found") + } + log.WithContext(ctx).Errorf("failed to get user invite from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get user invite from store") + } + + if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt invite: %w", err) + } + + return &invite, nil +} + +// GetUserInviteByEmail retrieves a user invite by account ID and email. +// Since email is encrypted with random IVs, we fetch all invites for the account +// and compare emails in memory after decryption. +func (s *SqlStore) GetUserInviteByEmail(ctx context.Context, lockStrength LockingStrength, accountID, email string) (*types.UserInviteRecord, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var invites []*types.UserInviteRecord + result := tx.Find(&invites, "account_id = ?", accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get user invites from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get user invites from store") + } + + for _, invite := range invites { + if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt invite: %w", err) + } + if strings.EqualFold(invite.Email, email) { + return invite, nil + } + } + + return nil, status.Errorf(status.NotFound, "user invite not found for email") +} + +// GetAccountUserInvites retrieves all user invites for an account +func (s *SqlStore) GetAccountUserInvites(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.UserInviteRecord, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var invites []*types.UserInviteRecord + result := tx.Find(&invites, "account_id = ?", accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get user invites from store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get user invites from store") + } + + for _, invite := range invites { + if err := invite.DecryptSensitiveData(s.fieldEncrypt); err != nil { + return nil, fmt.Errorf("decrypt invite: %w", err) + } + } + + return invites, nil +} + +// DeleteUserInvite deletes a user invite by its ID +func (s *SqlStore) DeleteUserInvite(ctx context.Context, inviteID string) error { + result := s.db.Delete(&types.UserInviteRecord{}, idQueryCondition, inviteID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete user invite from store: %s", result.Error) + return status.Errorf(status.Internal, "failed to delete user invite from store") + } + return nil +} diff --git a/management/server/store/sql_store_user_test.go b/management/server/store/sql_store_user_test.go new file mode 100644 index 000000000..34bc458fe --- /dev/null +++ b/management/server/store/sql_store_user_test.go @@ -0,0 +1,343 @@ +package store + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/management/server/util" + "github.com/netbirdio/netbird/shared/management/status" + "github.com/netbirdio/netbird/util/crypt" +) + +func TestSqlStore_GetAccountUsers(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + if err != nil { + t.Fatal(err) + } + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + account, err := store.GetAccount(context.Background(), accountID) + require.NoError(t, err) + users, err := store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Len(t, users, len(account.Users)) +} + +func TestSqlStore_GetUserByUserID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + tests := []struct { + name string + userID string + expectError bool + }{ + { + name: "retrieve existing user", + userID: "edafee4e-63fb-11ec-90d6-0242ac120003", + expectError: false, + }, + { + name: "retrieve non-existing user", + userID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty user ID", + userID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + user, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, tt.userID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, user) + } else { + require.NoError(t, err) + require.NotNil(t, user) + require.Equal(t, tt.userID, user.Id) + } + }) + } +} + +func TestSqlStore_GetUserByPATID(t *testing.T) { + store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "../testdata/store.sql", t.TempDir()) + t.Cleanup(cleanUp) + assert.NoError(t, err) + + id := "9dj38s35-63fb-11ec-90d6-0242ac120003" + + user, err := store.GetUserByPATID(context.Background(), LockingStrengthNone, id) + require.NoError(t, err) + require.Equal(t, "f4f6d672-63fb-11ec-90d6-0242ac120003", user.Id) +} + +func TestSqlStore_SaveUser(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + user := &types.User{ + Id: "user-id", + AccountID: accountID, + Role: types.UserRoleAdmin, + IsServiceUser: false, + AutoGroups: []string{"groupA", "groupB"}, + Blocked: false, + LastLogin: util.ToPtr(time.Now().UTC()), + CreatedAt: time.Now().UTC().Add(-time.Hour), + Issued: types.UserIssuedIntegration, + } + err = store.SaveUser(context.Background(), user) + require.NoError(t, err) + + saveUser, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, user.Id) + require.NoError(t, err) + require.Equal(t, user.Id, saveUser.Id) + require.Equal(t, user.AccountID, saveUser.AccountID) + require.Equal(t, user.Role, saveUser.Role) + require.Equal(t, user.AutoGroups, saveUser.AutoGroups) + require.WithinDurationf(t, user.GetLastLogin(), saveUser.LastLogin.UTC(), time.Millisecond, "LastLogin should be equal") + require.WithinDurationf(t, user.CreatedAt, saveUser.CreatedAt.UTC(), time.Millisecond, "CreatedAt should be equal") + require.Equal(t, user.Issued, saveUser.Issued) + require.Equal(t, user.Blocked, saveUser.Blocked) + require.Equal(t, user.IsServiceUser, saveUser.IsServiceUser) +} + +func TestSqlStore_SaveUsers(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + accountUsers, err := store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Len(t, accountUsers, 2) + + users := []*types.User{ + { + Id: "user-1", + AccountID: accountID, + Issued: "api", + AutoGroups: []string{"groupA", "groupB"}, + }, + { + Id: "user-2", + AccountID: accountID, + Issued: "integration", + AutoGroups: []string{"groupA"}, + }, + } + err = store.SaveUsers(context.Background(), users) + require.NoError(t, err) + + accountUsers, err = store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Len(t, accountUsers, 4) + + users[1].AutoGroups = []string{"groupA", "groupC"} + err = store.SaveUsers(context.Background(), users) + require.NoError(t, err) + + user, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, users[1].Id) + require.NoError(t, err) + require.Equal(t, users[1].AutoGroups, user.AutoGroups) +} + +func TestSqlStore_SaveUserWithEncryption(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + // Enable encryption + key, err := crypt.GenerateKey() + require.NoError(t, err) + fieldEncrypt, err := crypt.NewFieldEncrypt(key) + require.NoError(t, err) + store.SetFieldEncrypt(fieldEncrypt) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + // rawUser is used to read raw (potentially encrypted) data from the database + // without any gorm hooks or automatic decryption + type rawUser struct { + Id string + Email string + Name string + } + + t.Run("save user with empty email and name", func(t *testing.T) { + user := &types.User{ + Id: "user-empty-fields", + AccountID: accountID, + Role: types.UserRoleUser, + Email: "", + Name: "", + AutoGroups: []string{"groupA"}, + } + err = store.SaveUser(context.Background(), user) + require.NoError(t, err) + + // Verify using direct database query that empty strings remain empty (not encrypted) + var raw rawUser + err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", user.Id).First(&raw).Error + require.NoError(t, err) + require.Equal(t, "", raw.Email, "empty email should remain empty in database") + require.Equal(t, "", raw.Name, "empty name should remain empty in database") + + // Verify manual decryption returns empty strings + decryptedEmail, err := fieldEncrypt.Decrypt(raw.Email) + require.NoError(t, err) + require.Equal(t, "", decryptedEmail) + + decryptedName, err := fieldEncrypt.Decrypt(raw.Name) + require.NoError(t, err) + require.Equal(t, "", decryptedName) + }) + + t.Run("save user with email and name", func(t *testing.T) { + user := &types.User{ + Id: "user-with-fields", + AccountID: accountID, + Role: types.UserRoleAdmin, + Email: "test@example.com", + Name: "Test User", + AutoGroups: []string{"groupB"}, + } + err = store.SaveUser(context.Background(), user) + require.NoError(t, err) + + // Verify using direct database query that the data is encrypted (not plaintext) + var raw rawUser + err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", user.Id).First(&raw).Error + require.NoError(t, err) + require.NotEqual(t, "test@example.com", raw.Email, "email should be encrypted in database") + require.NotEqual(t, "Test User", raw.Name, "name should be encrypted in database") + + // Verify manual decryption returns correct values + decryptedEmail, err := fieldEncrypt.Decrypt(raw.Email) + require.NoError(t, err) + require.Equal(t, "test@example.com", decryptedEmail) + + decryptedName, err := fieldEncrypt.Decrypt(raw.Name) + require.NoError(t, err) + require.Equal(t, "Test User", decryptedName) + }) + + t.Run("save multiple users with mixed fields", func(t *testing.T) { + users := []*types.User{ + { + Id: "batch-user-1", + AccountID: accountID, + Email: "", + Name: "", + }, + { + Id: "batch-user-2", + AccountID: accountID, + Email: "batch@example.com", + Name: "Batch User", + }, + } + err = store.SaveUsers(context.Background(), users) + require.NoError(t, err) + + // Verify first user (empty fields) using direct database query + var raw1 rawUser + err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", "batch-user-1").First(&raw1).Error + require.NoError(t, err) + require.Equal(t, "", raw1.Email, "empty email should remain empty in database") + require.Equal(t, "", raw1.Name, "empty name should remain empty in database") + + // Verify second user (with fields) using direct database query + var raw2 rawUser + err = store.(*SqlStore).db.Table("users").Select("id, email, name").Where("id = ?", "batch-user-2").First(&raw2).Error + require.NoError(t, err) + require.NotEqual(t, "batch@example.com", raw2.Email, "email should be encrypted in database") + require.NotEqual(t, "Batch User", raw2.Name, "name should be encrypted in database") + + // Verify manual decryption returns empty strings for first user + decryptedEmail1, err := fieldEncrypt.Decrypt(raw1.Email) + require.NoError(t, err) + require.Equal(t, "", decryptedEmail1) + + decryptedName1, err := fieldEncrypt.Decrypt(raw1.Name) + require.NoError(t, err) + require.Equal(t, "", decryptedName1) + + // Verify manual decryption returns correct values for second user + decryptedEmail2, err := fieldEncrypt.Decrypt(raw2.Email) + require.NoError(t, err) + require.Equal(t, "batch@example.com", decryptedEmail2) + + decryptedName2, err := fieldEncrypt.Decrypt(raw2.Name) + require.NoError(t, err) + require.Equal(t, "Batch User", decryptedName2) + }) +} + +func TestSqlStore_DeleteUser(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + userID := "f4f6d672-63fb-11ec-90d6-0242ac120003" + + err = store.DeleteUser(context.Background(), accountID, userID) + require.NoError(t, err) + + user, err := store.GetUserByUserID(context.Background(), LockingStrengthNone, userID) + require.Error(t, err) + require.Nil(t, user) + + userPATs, err := store.GetUserPATs(context.Background(), LockingStrengthNone, userID) + require.NoError(t, err) + require.Len(t, userPATs, 0) +} + +func TestSqlStore_SaveUsers_LargeBatch(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + accountUsers, err := store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Len(t, accountUsers, 2) + + usersToSave := make([]*types.User, 0) + + for i := 1; i <= 8000; i++ { + usersToSave = append(usersToSave, &types.User{ + Id: fmt.Sprintf("user-%d", i), + AccountID: accountID, + Role: types.UserRoleUser, + }) + } + + err = store.SaveUsers(context.Background(), usersToSave) + require.NoError(t, err) + + accountUsers, err = store.GetAccountUsers(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.Equal(t, 8002, len(accountUsers)) +} diff --git a/management/server/store/sql_store_zone.go b/management/server/store/sql_store_zone.go new file mode 100644 index 000000000..02b65546a --- /dev/null +++ b/management/server/store/sql_store_zone.go @@ -0,0 +1,98 @@ +package store + +import ( + "context" + "errors" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/shared/management/status" +) + +func (s *SqlStore) CreateZone(ctx context.Context, zone *zones.Zone) error { + result := s.db.Create(zone) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to create zone to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to create zone to store") + } + + return nil +} + +func (s *SqlStore) UpdateZone(ctx context.Context, zone *zones.Zone) error { + result := s.db.Select("*").Save(zone) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to update zone to store: %v", result.Error) + return status.Errorf(status.Internal, "failed to update zone to store") + } + + return nil +} + +func (s *SqlStore) DeleteZone(ctx context.Context, accountID, zoneID string) error { + result := s.db.Delete(&zones.Zone{}, accountAndIDQueryCondition, accountID, zoneID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to delete zone from store: %v", result.Error) + return status.Errorf(status.Internal, "failed to delete zone from store") + } + + if result.RowsAffected == 0 { + return status.NewZoneNotFoundError(zoneID) + } + + return nil +} + +func (s *SqlStore) GetZoneByID(ctx context.Context, lockStrength LockingStrength, accountID, zoneID string) (*zones.Zone, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var zone *zones.Zone + result := tx.Preload("Records").Take(&zone, accountAndIDQueryCondition, accountID, zoneID) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewZoneNotFoundError(zoneID) + } + + log.WithContext(ctx).Errorf("failed to get zone from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get zone from store") + } + + return zone, nil +} + +func (s *SqlStore) GetZoneByDomain(ctx context.Context, accountID, domain string) (*zones.Zone, error) { + var zone *zones.Zone + result := s.db.Where("account_id = ? AND domain = ?", accountID, domain).First(&zone) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.NewZoneNotFoundError(domain) + } + + log.WithContext(ctx).Errorf("failed to get zone by domain from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get zone by domain from store") + } + + return zone, nil +} + +func (s *SqlStore) GetAccountZones(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*zones.Zone, error) { + tx := s.db + if lockStrength != LockingStrengthNone { + tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)}) + } + + var zones []*zones.Zone + result := tx.Preload("Records").Find(&zones, accountIDCondition, accountID) + if result.Error != nil { + log.WithContext(ctx).Errorf("failed to get zones from the store: %s", result.Error) + return nil, status.Errorf(status.Internal, "failed to get zones from store") + } + + return zones, nil +} diff --git a/management/server/store/sql_store_zone_test.go b/management/server/store/sql_store_zone_test.go new file mode 100644 index 000000000..da2ef12e6 --- /dev/null +++ b/management/server/store/sql_store_zone_test.go @@ -0,0 +1,238 @@ +package store + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestSqlStore_CreateZone(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + savedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, accountID, zone.ID) + require.NoError(t, err) + require.NotNil(t, savedZone) + assert.Equal(t, zone.ID, savedZone.ID) + assert.Equal(t, zone.Name, savedZone.Name) + assert.Equal(t, zone.Domain, savedZone.Domain) + assert.Equal(t, zone.Enabled, savedZone.Enabled) + assert.Equal(t, zone.EnableSearchDomain, savedZone.EnableSearchDomain) + assert.Equal(t, zone.DistributionGroups, savedZone.DistributionGroups) +} + +func TestSqlStore_GetZoneByID(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + zoneID string + expectError bool + }{ + { + name: "retrieve existing zone", + accountID: accountID, + zoneID: zone.ID, + expectError: false, + }, + { + name: "retrieve non-existing zone", + accountID: accountID, + zoneID: "non-existing", + expectError: true, + }, + { + name: "retrieve with empty zone ID", + accountID: accountID, + zoneID: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + savedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, tt.accountID, tt.zoneID) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) + require.Nil(t, savedZone) + } else { + require.NoError(t, err) + require.NotNil(t, savedZone) + assert.Equal(t, tt.zoneID, savedZone.ID) + } + }) + } +} + +func TestSqlStore_GetAccountZones(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone1 := zones.NewZone(accountID, "Zone 1", "example1.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone1) + require.NoError(t, err) + + zone2 := zones.NewZone(accountID, "Zone 2", "example2.com", true, true, []string{"group1", "group2"}) + err = store.CreateZone(context.Background(), zone2) + require.NoError(t, err) + + allZones, err := store.GetAccountZones(context.Background(), LockingStrengthNone, accountID) + require.NoError(t, err) + require.NotNil(t, allZones) + assert.GreaterOrEqual(t, len(allZones), 2) + + zoneIDs := make(map[string]bool) + for _, z := range allZones { + zoneIDs[z.ID] = true + } + assert.True(t, zoneIDs[zone1.ID]) + assert.True(t, zoneIDs[zone2.ID]) +} + +func TestSqlStore_GetZoneByDomain(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + otherAccountID := "bf1c8084-ba50-4ce7-9439-34653001fc3c" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + tests := []struct { + name string + accountID string + domain string + expectError bool + errorType status.Type + }{ + { + name: "retrieve existing zone by domain", + accountID: accountID, + domain: "example.com", + expectError: false, + }, + { + name: "retrieve non-existing zone domain", + accountID: accountID, + domain: "non-existing.com", + expectError: true, + errorType: status.NotFound, + }, + { + name: "retrieve with empty domain", + accountID: accountID, + domain: "", + expectError: true, + errorType: status.NotFound, + }, + { + name: "retrieve with different account ID", + accountID: otherAccountID, + domain: "example.com", + expectError: true, + errorType: status.NotFound, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + savedZone, err := store.GetZoneByDomain(context.Background(), tt.accountID, tt.domain) + if tt.expectError { + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, tt.errorType, sErr.Type()) + require.Nil(t, savedZone) + } else { + require.NoError(t, err) + require.NotNil(t, savedZone) + assert.Equal(t, tt.domain, savedZone.Domain) + assert.Equal(t, zone.ID, savedZone.ID) + assert.Equal(t, zone.Name, savedZone.Name) + } + }) + } +} + +func TestSqlStore_UpdateZone(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + zone.Name = "Updated Zone" + zone.Domain = "updated.com" + zone.Enabled = false + zone.EnableSearchDomain = true + zone.DistributionGroups = []string{"group2", "group3"} + + err = store.UpdateZone(context.Background(), zone) + require.NoError(t, err) + + updatedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, accountID, zone.ID) + require.NoError(t, err) + require.NotNil(t, updatedZone) + assert.Equal(t, "Updated Zone", updatedZone.Name) + assert.Equal(t, "updated.com", updatedZone.Domain) + assert.False(t, updatedZone.Enabled) + assert.True(t, updatedZone.EnableSearchDomain) + assert.Equal(t, []string{"group2", "group3"}, updatedZone.DistributionGroups) +} + +func TestSqlStore_DeleteZone(t *testing.T) { + store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/extended-store.sql", t.TempDir()) + t.Cleanup(cleanup) + require.NoError(t, err) + + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + + zone := zones.NewZone(accountID, "Test Zone", "example.com", true, false, []string{"group1"}) + err = store.CreateZone(context.Background(), zone) + require.NoError(t, err) + + err = store.DeleteZone(context.Background(), accountID, zone.ID) + require.NoError(t, err) + + deletedZone, err := store.GetZoneByID(context.Background(), LockingStrengthNone, accountID, zone.ID) + require.Error(t, err) + require.Nil(t, deletedZone) + sErr, ok := status.FromError(err) + require.True(t, ok) + require.Equal(t, sErr.Type(), status.NotFound) +} diff --git a/management/server/store/sqlstore_bench_test.go b/management/server/store/sqlstore_bench_test.go index a38b4a8c1..b1933b0d3 100644 --- a/management/server/store/sqlstore_bench_test.go +++ b/management/server/store/sqlstore_bench_test.go @@ -22,6 +22,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + nbdb "github.com/netbirdio/netbird/management/internals/shared/db" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" @@ -281,10 +282,11 @@ func setupBenchmarkDB(b testing.TB) (*SqlStore, func(), string) { b.Fatalf("failed to migrate database: %v", err) } - store := &SqlStore{ - db: db, - pool: pool, + conn, err := nbdb.NewConn(context.Background(), db, nbdb.PostgresStoreEngine, pool) + if err != nil { + b.Fatalf("failed to create connection: %v", err) } + store := &SqlStore{conn: conn, db: conn.DB(nil)} const ( accountID = "benchmark-account-id" @@ -527,7 +529,7 @@ func BenchmarkGetAccount(b *testing.B) { } } }) - store.pool.Close() + _ = store.Close(ctx) } func TestAccountEquivalence(t *testing.T) { diff --git a/management/server/store/store.go b/management/server/store/store.go index 97da95b4c..01aaf4892 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -13,7 +13,6 @@ import ( "path" "path/filepath" "regexp" - "runtime" "slices" "strings" "sync" @@ -23,19 +22,19 @@ import ( log "github.com/sirupsen/logrus" "gorm.io/driver/mysql" "gorm.io/driver/postgres" - "gorm.io/driver/sqlite" "gorm.io/gorm" "github.com/netbirdio/netbird/dns" - "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/internals/modules/zones" "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/internals/shared/db" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/testutil" "github.com/netbirdio/netbird/management/server/types" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/util" "github.com/netbirdio/netbird/util/crypt" @@ -49,14 +48,14 @@ import ( "github.com/netbirdio/netbird/route" ) -type LockingStrength string +type LockingStrength = db.LockingStrength const ( - LockingStrengthUpdate LockingStrength = "UPDATE" // Strongest lock, preventing any changes by other transactions until your transaction completes. - LockingStrengthShare LockingStrength = "SHARE" // Allows reading but prevents changes by other transactions. - LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE" // Similar to UPDATE but allows changes to related rows. - LockingStrengthKeyShare LockingStrength = "KEY SHARE" // Protects against changes to primary/unique keys but allows other updates. - LockingStrengthNone LockingStrength = "NONE" // No locking, allowing all transactions to proceed without restrictions. + LockingStrengthUpdate = db.LockingStrengthUpdate + LockingStrengthShare = db.LockingStrengthShare + LockingStrengthNoKeyUpdate = db.LockingStrengthNoKeyUpdate + LockingStrengthKeyShare = db.LockingStrengthKeyShare + LockingStrengthNone = db.LockingStrengthNone ) type Store interface { @@ -138,6 +137,7 @@ type Store interface { GetAccountPolicies(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.Policy, error) GetPolicyByID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*types.Policy, error) + GetPolicyByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*types.Policy, error) CreatePolicy(ctx context.Context, policy *types.Policy) error SavePolicy(ctx context.Context, policy *types.Policy) error DeletePolicy(ctx context.Context, accountID, policyID string) error @@ -160,7 +160,7 @@ type Store interface { RemoveResourceFromGroup(ctx context.Context, accountId string, groupID string, resourceID string) error AddPeerToAccount(ctx context.Context, peer *nbpeer.Peer) error GetPeerByPeerPubKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (*nbpeer.Peer, error) - GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) + GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) GetUserPeers(ctx context.Context, lockStrength LockingStrength, accountID, userID string) ([]*nbpeer.Peer, error) GetPeerByID(ctx context.Context, lockStrength LockingStrength, accountID string, peerID string) (*nbpeer.Peer, error) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error) @@ -208,6 +208,7 @@ type Store interface { GetAccountRoutes(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*route.Route, error) GetRouteByID(ctx context.Context, lockStrength LockingStrength, accountID, routeID string) (*route.Route, error) + GetRouteByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, routeID string) (*route.Route, error) SaveRoute(ctx context.Context, route *route.Route) error DeleteRoute(ctx context.Context, accountID, routeID string) error @@ -248,6 +249,7 @@ type Store interface { GetNetworkResourcesByNetID(ctx context.Context, lockStrength LockingStrength, accountID, netID string) ([]*resourceTypes.NetworkResource, error) GetNetworkResourcesByAccountID(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*resourceTypes.NetworkResource, error) GetNetworkResourceByID(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) (*resourceTypes.NetworkResource, error) + GetNetworkResourceByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) (*resourceTypes.NetworkResource, error) GetNetworkResourceByName(ctx context.Context, lockStrength LockingStrength, accountID, resourceName string) (*resourceTypes.NetworkResource, error) SaveNetworkResource(ctx context.Context, resource *resourceTypes.NetworkResource) error DeleteNetworkResource(ctx context.Context, accountID, resourceID string) error @@ -304,6 +306,7 @@ type Store interface { GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error) ListFreeDomains(ctx context.Context, accountID string) ([]string, error) ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) + LockCustomDomains(ctx context.Context, accountID string, serviceDomain nbdomain.Domain) ([]*domain.Domain, error) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) @@ -311,15 +314,13 @@ type Store interface { DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error - CreateAccessLog(ctx context.Context, log *accesslogs.AccessLogEntry) error - GetAccountAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) - DeleteOldAccessLogs(ctx context.Context, olderThan time.Time) (int64, error) CreateAgentNetworkAccessLog(ctx context.Context, entry *agentNetworkTypes.AgentNetworkAccessLog, groups []agentNetworkTypes.AgentNetworkAccessLogGroup) error CreateAgentNetworkUsage(ctx context.Context, usage *agentNetworkTypes.AgentNetworkUsage, groups []agentNetworkTypes.AgentNetworkUsageGroup) error GetAgentNetworkAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLog, int64, error) GetAgentNetworkAccessLogSessions(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLogSession, int64, error) GetAgentNetworkUsageRows(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkUsage, error) DeleteOldAgentNetworkAccessLogs(ctx context.Context, accountID string, olderThan time.Time) (int64, error) + GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx context.Context) ([]string, error) GetServiceTargetByTargetID(ctx context.Context, lockStrength LockingStrength, accountID string, targetID string) (*rpservice.Target, error) GetTargetsByServiceID(ctx context.Context, lockStrength LockingStrength, accountID string, serviceID string) ([]*rpservice.Target, error) DeleteTarget(ctx context.Context, accountID string, serviceID string, targetID uint) error @@ -335,6 +336,8 @@ type Store interface { GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool + GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error) DisconnectAllProxies(ctx context.Context) (int64, error) @@ -387,6 +390,7 @@ type Store interface { GetAgentNetworkConsumption(ctx context.Context, lockStrength LockingStrength, accountID string, kind agentNetworkTypes.ConsumptionDimension, dimID string, windowSeconds int64, windowStart time.Time) (*agentNetworkTypes.Consumption, error) GetAgentNetworkConsumptionBatch(ctx context.Context, lockStrength LockingStrength, accountID string, keys []agentNetworkTypes.ConsumptionKey) (map[agentNetworkTypes.ConsumptionKey]*agentNetworkTypes.Consumption, error) ListAgentNetworkConsumption(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*agentNetworkTypes.Consumption, error) + DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx context.Context) (int64, error) GetAccountAgentNetworkBudgetRules(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*agentNetworkTypes.AccountBudgetRule, error) GetAgentNetworkBudgetRuleByID(ctx context.Context, lockStrength LockingStrength, accountID, ruleID string) (*agentNetworkTypes.AccountBudgetRule, error) SaveAgentNetworkBudgetRule(ctx context.Context, rule *agentNetworkTypes.AccountBudgetRule) error @@ -487,7 +491,7 @@ func getStoreEngine(ctx context.Context, dataDir string, kind types.Engine) type // Migrate if it is the first run with a JSON file existing and no SQLite file present jsonStoreFile := filepath.Join(dataDir, storeFileName) - sqliteStoreFile := filepath.Join(dataDir, storeSqliteFileName) + sqliteStoreFile := filepath.Join(dataDir, db.SqliteFileName) if util.FileExists(jsonStoreFile) && !util.FileExists(sqliteStoreFile) { log.WithContext(ctx).Warnf("unsupported store engine specified, but found %s. Automatically migrating to SQLite.", jsonStoreFile) @@ -506,6 +510,16 @@ func getStoreEngine(ctx context.Context, dataDir string, kind types.Engine) type // NewStore creates a new store based on the provided engine type, data directory, and telemetry metrics func NewStore(ctx context.Context, kind types.Engine, dataDir string, metrics telemetry.AppMetrics, skipMigration bool) (Store, error) { + conn, err := OpenConn(ctx, kind, dataDir) + if err != nil { + return nil, err + } + return newStore(ctx, conn, metrics, skipMigration) +} + +// OpenConn resolves the configured engine and opens the connection that the +// store and the domain repositories share. +func OpenConn(ctx context.Context, kind types.Engine, dataDir string) (*db.Conn, error) { kind = getStoreEngine(ctx, dataDir, kind) if err := checkFileStoreEngine(kind, dataDir); err != nil { @@ -515,13 +529,21 @@ func NewStore(ctx context.Context, kind types.Engine, dataDir string, metrics te switch kind { case types.SqliteStoreEngine: log.WithContext(ctx).Info("using SQLite store engine") - return NewSqliteStore(ctx, dataDir, metrics, skipMigration) + return db.OpenSqlite(ctx, dataDir) case types.PostgresStoreEngine: log.WithContext(ctx).Info("using Postgres store engine") - return newPostgresStore(ctx, metrics, skipMigration) + dsn, ok := lookupDSNEnv(PostgresDsnEnv, PostgresDsnEnvLegacy) + if !ok { + return nil, fmt.Errorf("%s is not set", PostgresDsnEnv) + } + return db.OpenPostgres(ctx, dsn, db.DefaultPoolConfig) case types.MysqlStoreEngine: log.WithContext(ctx).Info("using MySQL store engine") - return newMysqlStore(ctx, metrics, skipMigration) + dsn, ok := lookupDSNEnv(mysqlDsnEnv, mysqlDsnEnvLegacy) + if !ok { + return nil, fmt.Errorf("%s is not set", mysqlDsnEnv) + } + return db.OpenMysql(ctx, dsn) default: return nil, fmt.Errorf("unsupported kind of store: %s", kind) } @@ -695,32 +717,28 @@ func NewTestStoreFromSQL(ctx context.Context, filename string, dataDir string) ( kind = types.SqliteStoreEngine } - storeStr := fmt.Sprintf("%s?cache=shared", storeSqliteFileName) - if runtime.GOOS == "windows" { - // Vo avoid `The process cannot access the file because it is being used by another process` on Windows - storeStr = storeSqliteFileName - } - - file := filepath.Join(dataDir, storeStr) - db, err := gorm.Open(sqlite.Open(file), getGormConfig()) + conn, err := db.OpenSqliteFile(ctx, dataDir, db.SqliteFileName) if err != nil { - return nil, nil, err + return nil, nil, fmt.Errorf("failed to create test store: %v", err) } if filename != "" { - err = LoadSQL(db, filename) + err = LoadSQL(conn.DB(nil), filename) if err != nil { + _ = conn.Close() return nil, nil, fmt.Errorf("failed to load SQL file: %v", err) } } - store, err := NewSqlStore(ctx, db, types.SqliteStoreEngine, nil, false) + store, err := NewSqlStore(ctx, conn, nil, false) if err != nil { + _ = conn.Close() return nil, nil, fmt.Errorf("failed to create test store: %v", err) } err = addAllGroupToAccount(ctx, store) if err != nil { + _ = store.Close(ctx) return nil, nil, fmt.Errorf("failed to add all group to account: %v", err) } @@ -785,9 +803,6 @@ func getSqlStoreEngine(ctx context.Context, sqliteStore *SqlStore, kind types.En closeConnection := func() { cleanup() store.Close(ctx) - if store.pool != nil { - store.pool.Close() - } if store != sqliteStore { // The sqlite store only seeded the engine under test; without this // every test leaks its connection and the opener goroutines. @@ -944,9 +959,6 @@ func postgresSchemaTemplate(ctx context.Context, baseDSN string, admin *gorm.DB) // TEMPLATE refuses a source that still has sessions, so release both handles // before the first clone. tplStore.Close(ctx) - if tplStore.pool != nil { - tplStore.pool.Close() - } schemaTemplates[key] = &schemaTemplate{dbName: name} return name, nil @@ -1031,11 +1043,11 @@ func mysqlTableNames(ctx context.Context, sqlDB *sql.DB, dbName string) ([]strin // cloneMysqlSchema replays the template's CREATE TABLE statements into the // database the DSN points at. func cloneMysqlSchema(ctx context.Context, dsn string, tableDDL []string) error { - db, err := gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), getGormConfig()) + gormDB, err := gorm.Open(mysql.Open(db.MysqlDSN(dsn)), db.GormConfig()) if err != nil { return fmt.Errorf("connect to test database: %w", err) } - sqlDB, err := db.DB() + sqlDB, err := gormDB.DB() if err != nil { return err } @@ -1075,19 +1087,19 @@ func closeGormDB(db *gorm.DB) { } func openDBWithRetry(dsn string, engine types.Engine, maxRetries int) (*gorm.DB, error) { - var db *gorm.DB + var gormDB *gorm.DB var err error for i := range maxRetries { switch engine { case types.PostgresStoreEngine: - db, err = gorm.Open(postgres.Open(dsn), &gorm.Config{}) + gormDB, err = gorm.Open(postgres.Open(dsn), &gorm.Config{}) case types.MysqlStoreEngine: - db, err = gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{}) + gormDB, err = gorm.Open(mysql.Open(db.MysqlDSN(dsn)), &gorm.Config{}) } if err == nil { - return db, nil + return gormDB, nil } if i < maxRetries-1 { @@ -1101,14 +1113,14 @@ func openDBWithRetry(dsn string, engine types.Engine, maxRetries int) (*gorm.DB, // createRandomDB creates a uniquely named database for one test. On postgres a // non-empty template is copied server-side with CREATE DATABASE ... TEMPLATE. -func createRandomDB(dsn string, db *gorm.DB, engine types.Engine, template string) (string, func(), error) { +func createRandomDB(dsn string, admin *gorm.DB, engine types.Engine, template string) (string, func(), error) { dbName := newTestDBName("test_db") createStmt := fmt.Sprintf("CREATE DATABASE %s", dbName) if template != "" && engine == types.PostgresStoreEngine { createStmt = fmt.Sprintf("CREATE DATABASE %s TEMPLATE %s", dbName, template) } - if err := execWithTemplateRetry(db, createStmt); err != nil { + if err := execWithTemplateRetry(admin, createStmt); err != nil { return "", nil, fmt.Errorf("failed to create database: %v", err) } @@ -1143,7 +1155,7 @@ func createRandomDB(dsn string, db *gorm.DB, engine types.Engine, template strin err = dropDB.Exec(fmt.Sprintf("DROP DATABASE IF EXISTS %s WITH (FORCE)", dbName)).Error case types.MysqlStoreEngine: - dropDB, err = gorm.Open(mysql.Open(originalDSN+"?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{ + dropDB, err = gorm.Open(mysql.Open(db.MysqlDSN(originalDSN)), &gorm.Config{ SkipDefaultTransaction: true, PrepareStmt: false, }) @@ -1221,7 +1233,7 @@ func MigrateFileStoreToSqlite(ctx context.Context, dataDir string) error { return fmt.Errorf("%s doesn't exist, couldn't continue the operation", fileStorePath) } - sqlStorePath := path.Join(dataDir, storeSqliteFileName) + sqlStorePath := path.Join(dataDir, db.SqliteFileName) if _, err := os.Stat(sqlStorePath); err == nil { return fmt.Errorf("%s already exists, couldn't continue the operation", sqlStorePath) } diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 399a07a19..956cac4b8 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -18,7 +18,6 @@ import ( dns "github.com/netbirdio/netbird/dns" types "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" - accesslogs "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" domain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" proxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" service "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" @@ -31,6 +30,7 @@ import ( posture "github.com/netbirdio/netbird/management/server/posture" types3 "github.com/netbirdio/netbird/management/server/types" route "github.com/netbirdio/netbird/route" + domain0 "github.com/netbirdio/netbird/shared/management/domain" crypt "github.com/netbirdio/netbird/util/crypt" gomock "go.uber.org/mock/gomock" ) @@ -246,20 +246,6 @@ func (mr *MockStoreMockRecorder) CountProxiesByAccountID(ctx, accountID any) *go return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountProxiesByAccountID", reflect.TypeOf((*MockStore)(nil).CountProxiesByAccountID), ctx, accountID) } -// CreateAccessLog mocks base method. -func (m *MockStore) CreateAccessLog(ctx context.Context, log *accesslogs.AccessLogEntry) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "CreateAccessLog", ctx, log) - ret0, _ := ret[0].(error) - return ret0 -} - -// CreateAccessLog indicates an expected call of CreateAccessLog. -func (mr *MockStoreMockRecorder) CreateAccessLog(ctx, log any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateAccessLog", reflect.TypeOf((*MockStore)(nil).CreateAccessLog), ctx, log) -} - // CreateAgentNetworkAccessLog mocks base method. func (m *MockStore) CreateAgentNetworkAccessLog(ctx context.Context, entry *types.AgentNetworkAccessLog, groups []types.AgentNetworkAccessLogGroup) error { m.ctrl.T.Helper() @@ -471,6 +457,21 @@ func (mr *MockStoreMockRecorder) DeleteAgentNetworkBudgetRule(ctx, accountID, ru return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAgentNetworkBudgetRule", reflect.TypeOf((*MockStore)(nil).DeleteAgentNetworkBudgetRule), ctx, accountID, ruleID) } +// DeleteAgentNetworkConsumptionOfDeletedAccounts mocks base method. +func (m *MockStore) DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx context.Context) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteAgentNetworkConsumptionOfDeletedAccounts", ctx) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// DeleteAgentNetworkConsumptionOfDeletedAccounts indicates an expected call of DeleteAgentNetworkConsumptionOfDeletedAccounts. +func (mr *MockStoreMockRecorder) DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAgentNetworkConsumptionOfDeletedAccounts", reflect.TypeOf((*MockStore)(nil).DeleteAgentNetworkConsumptionOfDeletedAccounts), ctx) +} + // DeleteAgentNetworkGuardrail mocks base method. func (m *MockStore) DeleteAgentNetworkGuardrail(ctx context.Context, accountID, guardrailID string) error { m.ctrl.T.Helper() @@ -668,21 +669,6 @@ func (mr *MockStoreMockRecorder) DeleteNetworkRouter(ctx, accountID, routerID an return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteNetworkRouter", reflect.TypeOf((*MockStore)(nil).DeleteNetworkRouter), ctx, accountID, routerID) } -// DeleteOldAccessLogs mocks base method. -func (m *MockStore) DeleteOldAccessLogs(ctx context.Context, olderThan time.Time) (int64, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DeleteOldAccessLogs", ctx, olderThan) - ret0, _ := ret[0].(int64) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// DeleteOldAccessLogs indicates an expected call of DeleteOldAccessLogs. -func (mr *MockStoreMockRecorder) DeleteOldAccessLogs(ctx, olderThan any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOldAccessLogs", reflect.TypeOf((*MockStore)(nil).DeleteOldAccessLogs), ctx, olderThan) -} - // DeleteOldAgentNetworkAccessLogs mocks base method. func (m *MockStore) DeleteOldAgentNetworkAccessLogs(ctx context.Context, accountID string, olderThan time.Time) (int64, error) { m.ctrl.T.Helper() @@ -967,22 +953,6 @@ func (mr *MockStoreMockRecorder) GetAccount(ctx, accountID any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccount", reflect.TypeOf((*MockStore)(nil).GetAccount), ctx, accountID) } -// GetAccountAccessLogs mocks base method. -func (m *MockStore) GetAccountAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetAccountAccessLogs", ctx, lockStrength, accountID, filter) - ret0, _ := ret[0].([]*accesslogs.AccessLogEntry) - ret1, _ := ret[1].(int64) - ret2, _ := ret[2].(error) - return ret0, ret1, ret2 -} - -// GetAccountAccessLogs indicates an expected call of GetAccountAccessLogs. -func (mr *MockStoreMockRecorder) GetAccountAccessLogs(ctx, lockStrength, accountID, filter any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountAccessLogs", reflect.TypeOf((*MockStore)(nil).GetAccountAccessLogs), ctx, lockStrength, accountID, filter) -} - // GetAccountAgentNetworkBudgetRules mocks base method. func (m *MockStore) GetAccountAgentNetworkBudgetRules(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.AccountBudgetRule, error) { m.ctrl.T.Helper() @@ -1360,18 +1330,18 @@ func (mr *MockStoreMockRecorder) GetAccountOwner(ctx, lockStrength, accountID an } // GetAccountPeers mocks base method. -func (m *MockStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*peer.Peer, error) { +func (m *MockStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetAccountPeers", ctx, lockStrength, accountID, nameFilter, ipFilter) + ret := m.ctrl.Call(m, "GetAccountPeers", ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter) ret0, _ := ret[0].([]*peer.Peer) ret1, _ := ret[1].(error) return ret0, ret1 } // GetAccountPeers indicates an expected call of GetAccountPeers. -func (mr *MockStoreMockRecorder) GetAccountPeers(ctx, lockStrength, accountID, nameFilter, ipFilter any) *gomock.Call { +func (mr *MockStoreMockRecorder) GetAccountPeers(ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockStore)(nil).GetAccountPeers), ctx, lockStrength, accountID, nameFilter, ipFilter) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockStore)(nil).GetAccountPeers), ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter) } // GetAccountPeersWithExpiration mocks base method. @@ -1584,6 +1554,21 @@ func (mr *MockStoreMockRecorder) GetActiveProxyClusterAddressesForAccount(ctx, a return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyClusterAddressesForAccount", reflect.TypeOf((*MockStore)(nil).GetActiveProxyClusterAddressesForAccount), ctx, accountID) } +// GetActiveProxyVersions mocks base method. +func (m *MockStore) GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetActiveProxyVersions", ctx, clusterAddr) + ret0, _ := ret[0].([]string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetActiveProxyVersions indicates an expected call of GetActiveProxyVersions. +func (mr *MockStoreMockRecorder) GetActiveProxyVersions(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyVersions", reflect.TypeOf((*MockStore)(nil).GetActiveProxyVersions), ctx, clusterAddr) +} + // GetAgentNetworkAccessLogSessions mocks base method. func (m *MockStore) GetAgentNetworkAccessLogSessions(ctx context.Context, lockStrength LockingStrength, accountID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) { m.ctrl.T.Helper() @@ -1885,6 +1870,20 @@ func (mr *MockStoreMockRecorder) GetAnyAccountID(ctx any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAnyAccountID", reflect.TypeOf((*MockStore)(nil).GetAnyAccountID), ctx) } +// GetClusterAllProxiesPrivate mocks base method. +func (m *MockStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetClusterAllProxiesPrivate", ctx, clusterAddr) + ret0, _ := ret[0].(*bool) + return ret0 +} + +// GetClusterAllProxiesPrivate indicates an expected call of GetClusterAllProxiesPrivate. +func (mr *MockStoreMockRecorder) GetClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusterAllProxiesPrivate", reflect.TypeOf((*MockStore)(nil).GetClusterAllProxiesPrivate), ctx, clusterAddr) +} + // GetClusterRequireSubdomain mocks base method. func (m *MockStore) GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { m.ctrl.T.Helper() @@ -2002,6 +2001,21 @@ func (mr *MockStoreMockRecorder) GetDNSRecordByID(ctx, lockStrength, accountID, return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDNSRecordByID", reflect.TypeOf((*MockStore)(nil).GetDNSRecordByID), ctx, lockStrength, accountID, zoneID, recordID) } +// GetDeletedAccountIDsWithAgentNetworkAccessLogs mocks base method. +func (m *MockStore) GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx context.Context) ([]string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetDeletedAccountIDsWithAgentNetworkAccessLogs", ctx) + ret0, _ := ret[0].([]string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetDeletedAccountIDsWithAgentNetworkAccessLogs indicates an expected call of GetDeletedAccountIDsWithAgentNetworkAccessLogs. +func (mr *MockStoreMockRecorder) GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetDeletedAccountIDsWithAgentNetworkAccessLogs", reflect.TypeOf((*MockStore)(nil).GetDeletedAccountIDsWithAgentNetworkAccessLogs), ctx) +} + // GetEmbeddedProxyPeerIDsByCluster mocks base method. func (m *MockStore) GetEmbeddedProxyPeerIDsByCluster(ctx context.Context, accountID string) (map[string][]string, error) { m.ctrl.T.Helper() @@ -2166,6 +2180,21 @@ func (mr *MockStoreMockRecorder) GetNetworkResourceByID(ctx, lockStrength, accou return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetNetworkResourceByID", reflect.TypeOf((*MockStore)(nil).GetNetworkResourceByID), ctx, lockStrength, accountID, resourceID) } +// GetNetworkResourceByIDOrPublicID mocks base method. +func (m *MockStore) GetNetworkResourceByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, resourceID string) (*types0.NetworkResource, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetNetworkResourceByIDOrPublicID", ctx, lockStrength, accountID, resourceID) + ret0, _ := ret[0].(*types0.NetworkResource) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetNetworkResourceByIDOrPublicID indicates an expected call of GetNetworkResourceByIDOrPublicID. +func (mr *MockStoreMockRecorder) GetNetworkResourceByIDOrPublicID(ctx, lockStrength, accountID, resourceID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetNetworkResourceByIDOrPublicID", reflect.TypeOf((*MockStore)(nil).GetNetworkResourceByIDOrPublicID), ctx, lockStrength, accountID, resourceID) +} + // GetNetworkResourceByName mocks base method. func (m *MockStore) GetNetworkResourceByName(ctx context.Context, lockStrength LockingStrength, accountID, resourceName string) (*types0.NetworkResource, error) { m.ctrl.T.Helper() @@ -2496,6 +2525,21 @@ func (mr *MockStoreMockRecorder) GetPolicyByID(ctx, lockStrength, accountID, pol return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPolicyByID", reflect.TypeOf((*MockStore)(nil).GetPolicyByID), ctx, lockStrength, accountID, policyID) } +// GetPolicyByIDOrPublicID mocks base method. +func (m *MockStore) GetPolicyByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*types3.Policy, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetPolicyByIDOrPublicID", ctx, lockStrength, accountID, policyID) + ret0, _ := ret[0].(*types3.Policy) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetPolicyByIDOrPublicID indicates an expected call of GetPolicyByIDOrPublicID. +func (mr *MockStoreMockRecorder) GetPolicyByIDOrPublicID(ctx, lockStrength, accountID, policyID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPolicyByIDOrPublicID", reflect.TypeOf((*MockStore)(nil).GetPolicyByIDOrPublicID), ctx, lockStrength, accountID, policyID) +} + // GetPolicyRulesByResourceID mocks base method. func (m *MockStore) GetPolicyRulesByResourceID(ctx context.Context, lockStrength LockingStrength, accountID, peerID string) ([]*types3.PolicyRule, error) { m.ctrl.T.Helper() @@ -2676,6 +2720,21 @@ func (mr *MockStoreMockRecorder) GetRouteByID(ctx, lockStrength, accountID, rout return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetRouteByID", reflect.TypeOf((*MockStore)(nil).GetRouteByID), ctx, lockStrength, accountID, routeID) } +// GetRouteByIDOrPublicID mocks base method. +func (m *MockStore) GetRouteByIDOrPublicID(ctx context.Context, lockStrength LockingStrength, accountID, routeID string) (*route.Route, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetRouteByIDOrPublicID", ctx, lockStrength, accountID, routeID) + ret0, _ := ret[0].(*route.Route) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetRouteByIDOrPublicID indicates an expected call of GetRouteByIDOrPublicID. +func (mr *MockStoreMockRecorder) GetRouteByIDOrPublicID(ctx, lockStrength, accountID, routeID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetRouteByIDOrPublicID", reflect.TypeOf((*MockStore)(nil).GetRouteByIDOrPublicID), ctx, lockStrength, accountID, routeID) +} + // GetRoutingPeerNetworks mocks base method. func (m *MockStore) GetRoutingPeerNetworks(ctx context.Context, accountID, peerID string) ([]string, error) { m.ctrl.T.Helper() @@ -3257,6 +3316,21 @@ func (mr *MockStoreMockRecorder) ListFreeDomains(ctx, accountID any) *gomock.Cal return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListFreeDomains", reflect.TypeOf((*MockStore)(nil).ListFreeDomains), ctx, accountID) } +// LockCustomDomains mocks base method. +func (m *MockStore) LockCustomDomains(ctx context.Context, accountID string, serviceDomain domain0.Domain) ([]*domain.Domain, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "LockCustomDomains", ctx, accountID, serviceDomain) + ret0, _ := ret[0].([]*domain.Domain) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// LockCustomDomains indicates an expected call of LockCustomDomains. +func (mr *MockStoreMockRecorder) LockCustomDomains(ctx, accountID, serviceDomain any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LockCustomDomains", reflect.TypeOf((*MockStore)(nil).LockCustomDomains), ctx, accountID, serviceDomain) +} + // MarkAccountPrimary mocks base method. func (m *MockStore) MarkAccountPrimary(ctx context.Context, accountID string) error { m.ctrl.T.Helper() diff --git a/management/server/types/networkmap_components_test.go b/management/server/types/networkmap_components_test.go index f6d542609..9e76de775 100644 --- a/management/server/types/networkmap_components_test.go +++ b/management/server/types/networkmap_components_test.go @@ -175,6 +175,63 @@ func TestNetworkMapComponents_NetworkResourceRoutes_RouterPeer(t *testing.T) { assert.NotEmpty(t, nm.RoutesFirewallRules, "router peer should have route firewall rules for the resource") } +// A receiver without a firewall asks Calculate to skip the route firewall +// rules. Everything the rest of the sync consumes — routes, peers, peer +// firewall rules — must come out unchanged. +func TestNetworkMapComponents_SkipRouteFirewallRules(t *testing.T) { + ctx := context.Background() + account := createComponentTestAccount() + + // The shared fixture leaves peer-router-1 out of every peer ACL, so its + // FirewallRules would be empty and the comparison below vacuous. Give the + // router a policy of its own. + account.Policies = append(account.Policies, &types.Policy{ + ID: "policy-router", Name: "Router connectivity", Enabled: true, + Rules: []*types.PolicyRule{{ + ID: "rule-router", Name: "Allow all <-> router", Enabled: true, + Action: types.PolicyTrafficActionAccept, Protocol: types.PolicyRuleProtocolALL, + Bidirectional: true, + Sources: []string{"group-all"}, Destinations: []string{"group-all"}, + }}, + }) + + validated := allPeersValidated(account) + + components := account.GetPeerNetworkMapComponents( + ctx, + "peer-router-1", + account.GetPeersCustomZone(ctx, "netbird.io"), + nil, + validated, + account.GetResourcePoliciesMap(), + account.GetResourceRoutersMap(), + account.GetActiveGroupUsers(), + ) + + full := components.Calculate(ctx) + require.NotEmpty(t, full.RoutesFirewallRules, "baseline: router peer must get route firewall rules") + require.NotEmpty(t, full.FirewallRules, "baseline: router peer must get peer firewall rules") + + components.SkipRouteFirewallRules = true + skipped := components.Calculate(ctx) + + assert.Empty(t, skipped.RoutesFirewallRules, "route firewall rules must not be computed when skipped") + assert.ElementsMatch(t, routeNetworks(full.Routes), routeNetworks(skipped.Routes), + "skipping route firewall rules must not change the routes") + assert.ElementsMatch(t, peerIDs(full.Peers), peerIDs(skipped.Peers), + "skipping route firewall rules must not change the peers to connect") + assert.Equal(t, full.FirewallRules, skipped.FirewallRules, + "peer firewall rules are unrelated and must come out unchanged") +} + +func routeNetworks(routes []*nmdata.Route) []string { + networks := make([]string, 0, len(routes)) + for _, r := range routes { + networks = append(networks, r.Network.String()) + } + return networks +} + func TestNetworkMapComponents_NetworkResourceRoutes_UnrelatedPeer(t *testing.T) { account := createComponentTestAccount() validated := allPeersValidated(account) diff --git a/management/server/types/policy.go b/management/server/types/policy.go index 0f7298d18..9786d17b6 100644 --- a/management/server/types/policy.go +++ b/management/server/types/policy.go @@ -29,7 +29,7 @@ type Policy struct { // ID of the policy' ID string `gorm:"primaryKey"` - PublicID string `json:"-"` + PublicID string `json:"-" gorm:"index"` // AccountID is a reference to Account that this object belongs AccountID string `json:"-" gorm:"index"` diff --git a/management/server/types/store.go b/management/server/types/store.go index 2ca4383b2..a13d52f91 100644 --- a/management/server/types/store.go +++ b/management/server/types/store.go @@ -1,10 +1,12 @@ package types -type Engine string +import "github.com/netbirdio/netbird/management/internals/shared/db" + +type Engine = db.Engine const ( - PostgresStoreEngine Engine = "postgres" + PostgresStoreEngine = db.PostgresStoreEngine FileStoreEngine Engine = "jsonfile" - SqliteStoreEngine Engine = "sqlite" - MysqlStoreEngine Engine = "mysql" + SqliteStoreEngine = db.SqliteStoreEngine + MysqlStoreEngine = db.MysqlStoreEngine ) diff --git a/management/server/user.go b/management/server/user.go index 3510a624b..5f29f4df7 100644 --- a/management/server/user.go +++ b/management/server/user.go @@ -861,9 +861,11 @@ func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transact allGroupChanges := slices.Concat(removedGroups, addedGroups) change.LinkGroups = allGroupChanges - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges) + if err != nil { return change, nil, nil, nil, fmt.Errorf("reconcile IPv6 for group changes: %w", err) } + change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...) } userEventsToAdd := am.prepareUserUpdateEvents(ctx, updatedUser.AccountID, initiatorUserId, oldUser, updatedUser, transferredOwnerRole, isNewUser, removedGroups, addedGroups, transaction) diff --git a/proxy/Dockerfile.ubi b/proxy/Dockerfile.ubi new file mode 100644 index 000000000..a74280a49 --- /dev/null +++ b/proxy/Dockerfile.ubi @@ -0,0 +1,32 @@ +FROM registry.access.redhat.com/ubi9/ubi-minimal@sha256:7fbeae18dc9476399f565e68255f602a3374ea8614ba3d14843565131a13ff93 + +ARG TARGETPLATFORM +ARG VERSION=dev +ARG RELEASE=1 + +LABEL name="netbird-reverse-proxy" \ + maintainer="NetBird " \ + vendor="NetBird GmbH" \ + version="${VERSION}" \ + release="${RELEASE}" \ + summary="NetBird Reverse Proxy" \ + description="NetBird Reverse Proxy provides an identity-aware entrypoint to services in NetBird networks." + +COPY --chmod=0555 ${TARGETPLATFORM}/netbird-proxy /go/bin/netbird-proxy +COPY licenses/ /licenses/ +# Only the writable directories share the root group for arbitrary non-root UIDs. +# Runtime-created private keys retain the application's restrictive file modes. +RUN mkdir -p /var/lib/netbird /certs && \ + chown 1000:0 /var/lib/netbird /certs && \ + chmod 0770 /var/lib/netbird /certs && \ + chmod -R a+rX /licenses + +USER 1000:0 +ENV HOME=/var/lib/netbird +ENV NB_PROXY_ADDRESS=":8443" +# Unprivileged ports: runtimes such as OpenShift and Podman keep the kernel +# default that reserves ports below 1024 for root. 8080 is the health probe. +ENV NB_PROXY_ACME_ADDRESS=":8081" +EXPOSE 8443 +STOPSIGNAL SIGTERM +ENTRYPOINT ["/go/bin/netbird-proxy"] diff --git a/proxy/auth/auth.go b/proxy/auth/auth.go index 5512bf003..605780959 100644 --- a/proxy/auth/auth.go +++ b/proxy/auth/auth.go @@ -30,6 +30,14 @@ const ( SessionJWTIssuer = "netbird-management" ) +// Query parameters management uses to hand the OIDC session to the proxy. The +// proxy strips them before forwarding, so they must not collide with names the +// proxied service uses itself. +const ( + SessionCodeQueryParam = "nb_session_code" + SessionTokenQueryParam = "session_token" +) + // HeaderUserID is the synthetic user id recorded for header-authenticated // requests. Header auth validates a per-service secret and resolves no user // record, so proxy access logs and management-minted session tokens both @@ -66,7 +74,7 @@ func ValidateSessionJWT(tokenString, domain string, publicKey ed25519.PublicKey) return nil, fmt.Errorf("unexpected signing method: %v", t.Header["alg"]) } return publicKey, nil - }, jwt.WithAudience(domain), jwt.WithIssuer(SessionJWTIssuer)) + }, jwt.WithAudience(domain), jwt.WithIssuer(SessionJWTIssuer), jwt.WithStrictDecoding()) if err != nil { return "", "", "", nil, nil, fmt.Errorf("parse token: %w", err) } diff --git a/proxy/cmd/proxy/cmd/root.go b/proxy/cmd/proxy/cmd/root.go index 9b180a5c4..765d5c05a 100644 --- a/proxy/cmd/proxy/cmd/root.go +++ b/proxy/cmd/proxy/cmd/root.go @@ -14,6 +14,7 @@ import ( "golang.org/x/crypto/acme" "github.com/netbirdio/netbird/shared/management/domain" + "github.com/netbirdio/netbird/shared/profiling" "github.com/netbirdio/netbird/client/embed" "github.com/netbirdio/netbird/proxy" @@ -30,6 +31,8 @@ const ( // how many buffers each receive/TUN worker eagerly allocates. Zero // (unset) keeps the platform default. envMaxBatchSize = "NB_PROXY_MAX_BATCH_SIZE" + + applicationName = "proxy" ) const DefaultManagementURL = "https://api.netbird.io:443" @@ -160,6 +163,9 @@ func runServer(cmd *cobra.Command, args []string) error { logger.Infof("configured log level: %s", level) + stopProfiling := profiling.Start(applicationName) + defer stopProfiling() + var wgPool, wgBatch uint64 var perf embed.Performance if raw := os.Getenv(envPreallocatedBuffers); raw != "" { diff --git a/proxy/cmd/proxy/main.go b/proxy/cmd/proxy/main.go index 16e7e8ac2..6851c6cfc 100644 --- a/proxy/cmd/proxy/main.go +++ b/proxy/cmd/proxy/main.go @@ -4,6 +4,7 @@ import ( "net/http" // nolint:gosec _ "net/http/pprof" + "os" "runtime" log "github.com/sirupsen/logrus" @@ -26,9 +27,13 @@ var ( ) func main() { - go func() { - log.Println(http.ListenAndServe("localhost:6060", nil)) - }() + if pprofAddr := os.Getenv("NB_PPROF_ADDR"); pprofAddr != "" { + log.Infof("pprof enabled, listening on: %s", pprofAddr) + go func() { + log.Println(http.ListenAndServe(pprofAddr, nil)) + }() + } + cmd.SetVersionInfo(Version, Commit, BuildDate, GoVersion) cmd.Execute() } diff --git a/proxy/collect-licenses.sh b/proxy/collect-licenses.sh new file mode 100644 index 000000000..ccf5f4dd6 --- /dev/null +++ b/proxy/collect-licenses.sh @@ -0,0 +1,82 @@ +#!/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-proxy-licenses.modules.XXXXXX") +sorted_modules=$(mktemp "${TMPDIR:-/tmp}/netbird-proxy-licenses.sorted.XXXXXX") + +if [ -e "$output" ] || [ -L "$output" ]; then + printf 'output directory already exists: %s\n' "$output" >&2 + exit 1 +fi +# Assemble beside the target and rename on success, so a failed run leaves +# nothing behind that would block the next attempt. +staging=$(mktemp -d "$output_parent/.$output_name.XXXXXX") +trap 'rm -f "$modules" "$sorted_modules"; rm -rf "$staging"' EXIT HUP INT TERM +mkdir "$staging/third_party" + +cp "$repo_root/proxy/LICENSE" "$staging/AGPL-3.0.txt" +cp "$repo_root/LICENSE" "$staging/BSD-3-Clause.txt" +node "$repo_root/proxy/web/scripts/third-party-licenses.mjs" >"$staging/Web-THIRD-PARTY-LICENSES" + +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}}' ./proxy/cmd/proxy >>"$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" "$staging/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="$staging/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" + +mv "$staging" "$output" diff --git a/proxy/internal/auth/README.md b/proxy/internal/auth/README.md new file mode 100644 index 000000000..5ebf75cee --- /dev/null +++ b/proxy/internal/auth/README.md @@ -0,0 +1,26 @@ +# PIN and password authentication limits + +PIN and password credentials are accepted only in a POST form body. Query-string +credentials and credentials on other HTTP methods are ignored. + +The proxy permits a burst of five credential checks per account and service, +then replenishes one check every six seconds (ten per minute). PIN and password +checks share the same budget. Five failed checks from one client IP in a +rolling five-minute window block that source for fifteen minutes. In-flight checks +reserve failure slots; blocked requests do not extend the cooldown. Successful +authentication clears that source's failure history. Infrastructure failures +consume the service budget without counting as incorrect credentials. + +Throttled requests return HTTP 429 with a `Retry-After` delay in seconds. The +login page displays that delay. Existing authenticated sessions and other +authentication methods do not consume these credential budgets. + +The client IP comes from the existing trusted-proxy resolution. Deployments +behind a load balancer must configure trusted proxies correctly; otherwise +visitors share the load balancer's source budget. Visitors behind the same NAT +also share a source budget for a service. + +State is held in memory per proxy process and resets on restart. Multiple +replicas have independent budgets. State is bounded to 16,384 source entries and +4,096 service entries; when capacity is exhausted, new checks are denied until +idle entries expire. Active blocks are never evicted to admit a new source. diff --git a/proxy/internal/auth/credential.go b/proxy/internal/auth/credential.go new file mode 100644 index 000000000..0e93891fb --- /dev/null +++ b/proxy/internal/auth/credential.go @@ -0,0 +1,100 @@ +package auth + +import ( + "errors" + "math" + "net/http" + "strconv" + "time" + + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/proxy/auth" + "github.com/netbirdio/netbird/proxy/internal/proxy" +) + +var errCredentialClientIP = errors.New("invalid client address") + +type credentialLimitError struct { + retryAfter time.Duration +} + +func (e *credentialLimitError) Error() string { + return "too many authentication attempts" +} + +func credentialFormValue(r *http.Request, field string) string { + if r.Method != http.MethodPost { + return "" + } + return r.PostFormValue(field) +} + +func (mw *Middleware) authenticateScheme(r *http.Request, config DomainConfig, scheme Scheme) (string, string, error) { + method := scheme.Type() + if (method != auth.MethodPIN && method != auth.MethodPassword) || !wasCredentialSubmitted(r, method) { + return scheme.Authenticate(r) + } + ip := mw.resolveClientIP(r).Unmap() + if !ip.IsValid() { + return "", "", errCredentialClientIP + } + source, retry := mw.credentials.begin(credentialSourceKey{ + service: credentialServiceKey{accountID: config.AccountID, serviceID: config.ServiceID}, + ip: ip, + }) + if retry > 0 { + return "", "", &credentialLimitError{retryAfter: retry} + } + token, prompt, err := scheme.Authenticate(r) + outcome := credentialUnavailable + if err == nil { + outcome = credentialRejected + if token != "" { + outcome = credentialAccepted + } + } + mw.credentials.finish(source, outcome) + return token, prompt, err +} + +func credentialRetryAfter(err error) time.Duration { + var limitErr *credentialLimitError + if errors.As(err, &limitErr) { + return limitErr.retryAfter + } + s := status.Convert(err) + if s.Code() != codes.ResourceExhausted { + return 0 + } + for _, detail := range s.Details() { + if info, ok := detail.(*errdetails.RetryInfo); ok && info.RetryDelay != nil && info.RetryDelay.CheckValid() == nil { + if delay := info.RetryDelay.AsDuration(); delay > 0 { + return delay + } + } + } + return credentialCheckInterval +} + +func (mw *Middleware) writeAuthenticationError(w http.ResponseWriter, r *http.Request, method auth.Method, err error) { + if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { + cd.SetOrigin(proxy.OriginAuth) + cd.SetAuthMethod(method.String()) + } + if retry := credentialRetryAfter(err); retry > 0 { + // RFC 6585 section 4 forbids caching 429 responses. + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Retry-After", strconv.FormatInt(int64(math.Ceil(retry.Seconds())), 10)) + http.Error(w, "too many authentication attempts; try again later", http.StatusTooManyRequests) + return + } + if errors.Is(err, errCredentialClientIP) { + http.Error(w, "invalid client address", http.StatusBadRequest) + return + } + mw.logger.WithField("scheme", method.String()).Warnf("authentication infrastructure error: %v", err) + http.Error(w, "authentication service unavailable", http.StatusBadGateway) +} diff --git a/proxy/internal/auth/credential_limiter.go b/proxy/internal/auth/credential_limiter.go new file mode 100644 index 000000000..be9e70413 --- /dev/null +++ b/proxy/internal/auth/credential_limiter.go @@ -0,0 +1,169 @@ +package auth + +import ( + "net/netip" + "sync" + "time" + + "golang.org/x/time/rate" + + "github.com/netbirdio/netbird/proxy/internal/types" +) + +const ( + credentialFailureLimit = 5 + credentialFailureWindow = 5 * time.Minute + credentialBlockDuration = 15 * time.Minute + credentialCheckInterval = 6 * time.Second + credentialCheckBurst = 5 + credentialMaxSources = 16384 + credentialMaxServices = 4096 + credentialCleanupInterval = time.Minute +) + +type credentialServiceKey struct { + accountID types.AccountID + serviceID types.ServiceID +} + +type credentialSourceKey struct { + service credentialServiceKey + ip netip.Addr +} + +type credentialSource struct { + failures []time.Time + pending int + expiresAt time.Time + blockedUntil time.Time +} + +type credentialService struct { + limiter *rate.Limiter + lastUsed time.Time +} + +type credentialOutcome string + +const ( + credentialUnavailable credentialOutcome = "unavailable" + credentialRejected credentialOutcome = "rejected" + credentialAccepted credentialOutcome = "accepted" +) + +// State is local to this proxy process. Active blocks are never evicted to +// make room for a new source; exhausting capacity denies new checks. +type credentialLimiter struct { + mu sync.Mutex + now func() time.Time + sources map[credentialSourceKey]*credentialSource + services map[credentialServiceKey]*credentialService + nextCleanup time.Time +} + +func newCredentialLimiter() *credentialLimiter { + return &credentialLimiter{ + now: time.Now, + sources: make(map[credentialSourceKey]*credentialSource), + services: make(map[credentialServiceKey]*credentialService), + } +} + +func (l *credentialLimiter) begin(key credentialSourceKey) (*credentialSource, time.Duration) { + l.mu.Lock() + defer l.mu.Unlock() + now := l.now() + l.cleanup(now) + source := l.sources[key] + if source != nil { + if now.Before(source.blockedUntil) { + return nil, source.blockedUntil.Sub(now) + } + if source.pending == 0 && !now.Before(source.expiresAt) { + *source = credentialSource{} + } + source.expireFailures(now) + // Reserve the failure budget before verification so concurrent guesses + // cannot all pass a check against the same completed failure count. + if len(source.failures)+source.pending >= credentialFailureLimit { + return nil, time.Second + } + } else if len(l.sources) >= credentialMaxSources { + return nil, credentialCleanupInterval + } + if retry := l.allowService(key.service, now); retry > 0 { + return nil, retry + } + if source == nil { + source = &credentialSource{} + l.sources[key] = source + } + if source.expiresAt.IsZero() { + source.expiresAt = now.Add(credentialFailureWindow) + } + source.pending++ + return source, 0 +} + +func (l *credentialLimiter) allowService(key credentialServiceKey, now time.Time) time.Duration { + service := l.services[key] + if service == nil { + if len(l.services) >= credentialMaxServices { + return credentialCleanupInterval + } + service = &credentialService{limiter: rate.NewLimiter(rate.Every(credentialCheckInterval), credentialCheckBurst)} + l.services[key] = service + } + service.lastUsed = now + if service.limiter.AllowN(now, 1) { + return 0 + } + return max(time.Nanosecond, time.Duration((1-service.limiter.TokensAt(now))*float64(credentialCheckInterval))) +} + +func (l *credentialLimiter) finish(source *credentialSource, outcome credentialOutcome) { + l.mu.Lock() + defer l.mu.Unlock() + source.pending-- + now := l.now() + source.expireFailures(now) + switch outcome { + case credentialRejected: + source.failures = append(source.failures, now) + source.expiresAt = now.Add(credentialFailureWindow) + if len(source.failures) >= credentialFailureLimit && source.blockedUntil.IsZero() { + source.blockedUntil = now.Add(credentialBlockDuration) + source.expiresAt = source.blockedUntil + } + case credentialAccepted: + if !now.Before(source.blockedUntil) { + source.failures = nil + source.expiresAt = now.Add(credentialFailureWindow) + } + case credentialUnavailable: + // Transport failures consume the service budget, but are not bad guesses. + } +} + +func (s *credentialSource) expireFailures(now time.Time) { + for len(s.failures) > 0 && !now.Before(s.failures[0].Add(credentialFailureWindow)) { + s.failures = s.failures[1:] + } +} + +func (l *credentialLimiter) cleanup(now time.Time) { + if now.Before(l.nextCleanup) { + return + } + l.nextCleanup = now.Add(credentialCleanupInterval) + for key, source := range l.sources { + if source.pending == 0 && !now.Before(source.expiresAt) { + delete(l.sources, key) + } + } + for key, service := range l.services { + if now.Sub(service.lastUsed) >= credentialBlockDuration { + delete(l.services, key) + } + } +} diff --git a/proxy/internal/auth/credential_limiter_test.go b/proxy/internal/auth/credential_limiter_test.go new file mode 100644 index 000000000..dfa346baf --- /dev/null +++ b/proxy/internal/auth/credential_limiter_test.go @@ -0,0 +1,190 @@ +package auth + +import ( + "net/netip" + "strconv" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/proxy/internal/types" +) + +func TestCredentialLimiterCooldown(t *testing.T) { + l := newCredentialLimiter() + now := time.Now() + l.now = func() time.Time { return now } + key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")} + for range credentialFailureLimit { + attempt, retry := l.begin(key) + require.Zero(t, retry, "initial guesses must reach verification") + l.finish(attempt, credentialRejected) + } + _, retry := l.begin(key) + assert.Equal(t, credentialBlockDuration, retry, "five failures must start a fifteen-minute block") + now = now.Add(credentialBlockDuration - time.Second) + _, retry = l.begin(key) + assert.Equal(t, time.Second, retry, "blocked requests must not extend the deadline") + now = now.Add(time.Second) + attempt, retry := l.begin(key) + require.Zero(t, retry, "the source must recover when its block expires") + l.finish(attempt, credentialAccepted) +} + +func TestCredentialLimiterFailureWindowAndSuccess(t *testing.T) { + for _, outcome := range []credentialOutcome{credentialAccepted, credentialUnavailable} { + t.Run(map[credentialOutcome]string{credentialAccepted: "success", credentialUnavailable: "infrastructure error"}[outcome], func(t *testing.T) { + l := newCredentialLimiter() + now := time.Now() + l.now = func() time.Time { return now } + key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")} + for range 4 { + attempt, retry := l.begin(key) + require.Zero(t, retry, "four failures must fit the budget") + l.finish(attempt, credentialRejected) + } + attempt, retry := l.begin(key) + require.Zero(t, retry, "fifth check must be allowed") + l.finish(attempt, outcome) + now = now.Add(credentialCheckInterval) + attempt, retry = l.begin(key) + require.Zero(t, retry, "success or infrastructure error must not start a block") + l.finish(attempt, credentialRejected) + now = now.Add(credentialCheckInterval) + attempt, retry = l.begin(key) + if outcome == credentialUnavailable { + assert.Greater(t, retry, time.Duration(0), "infrastructure errors must preserve earlier failures") + return + } + require.Zero(t, retry, "success must clear earlier failures") + l.finish(attempt, credentialRejected) + now = now.Add(credentialFailureWindow) + for range credentialFailureLimit { + attempt, retry = l.begin(key) + require.Zero(t, retry, "old failures must expire") + l.finish(attempt, credentialRejected) + } + }) + } +} + +func TestCredentialLimiterRollingWindow(t *testing.T) { + l := newCredentialLimiter() + now := time.Now() + l.now = func() time.Time { return now } + key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")} + attempt, retry := l.begin(key) + require.Zero(t, retry, "the first failure starts the history") + l.finish(attempt, credentialRejected) + now = now.Add(4 * time.Minute) + for range 3 { + attempt, retry = l.begin(key) + require.Zero(t, retry, "three more failures must fit the budget") + l.finish(attempt, credentialRejected) + } + now = now.Add(time.Minute + time.Second) + for range 2 { + attempt, retry = l.begin(key) + require.Zero(t, retry, "only the oldest failure must have expired") + l.finish(attempt, credentialRejected) + } + _, retry = l.begin(key) + assert.Equal(t, credentialBlockDuration, retry, "five recent failures must block even across the first window boundary") +} + +func TestCredentialLimiterServiceBudget(t *testing.T) { + l := newCredentialLimiter() + now := time.Now() + l.now = func() time.Time { return now } + key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")} + for range credentialCheckBurst { + attempt, retry := l.begin(key) + require.Zero(t, retry, "initial checks must fit the service burst") + l.finish(attempt, credentialAccepted) + key.ip = key.ip.Next() + } + _, retry := l.begin(key) + assert.Equal(t, credentialCheckInterval, retry, "changing IP must not bypass the service budget") + other := key + other.service.accountID = "another-account" + attempt, retry := l.begin(other) + require.Zero(t, retry, "accounts must have separate budgets") + l.finish(attempt, credentialAccepted) + other = key + other.service.serviceID = "another-service" + attempt, retry = l.begin(other) + require.Zero(t, retry, "services must have separate budgets") + l.finish(attempt, credentialAccepted) + now = now.Add(credentialCheckInterval) + attempt, retry = l.begin(key) + require.Zero(t, retry, "one check must refill every six seconds") + l.finish(attempt, credentialAccepted) + _, retry = l.begin(key) + assert.Equal(t, credentialCheckInterval, retry, "refill must only grant one new check") +} + +func TestCredentialLimiterConcurrentReservations(t *testing.T) { + l := newCredentialLimiter() + now := time.Now() + l.now = func() time.Time { return now } + key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")} + var attempts []*credentialSource + for range credentialFailureLimit { + attempt, retry := l.begin(key) + require.Zero(t, retry, "initial requests must reserve the failure budget") + attempts = append(attempts, attempt) + } + // Refill the service budget while earlier verification calls are still running. + now = now.Add(time.Minute) + var admitted atomic.Int32 + var wg sync.WaitGroup + for range 100 { + wg.Go(func() { + attempt, retry := l.begin(key) + if retry == 0 { + admitted.Add(1) + l.finish(attempt, credentialRejected) + } + }) + } + wg.Wait() + assert.Zero(t, admitted.Load(), "in-flight guesses must reserve the failure budget despite a refilled service budget") + for _, attempt := range attempts { + wg.Go(func() { l.finish(attempt, credentialRejected) }) + } + wg.Wait() + _, retry := l.begin(key) + assert.Equal(t, credentialBlockDuration, retry, "concurrent failures must activate the block") +} + +func TestCredentialLimiterCapacityAndCleanup(t *testing.T) { + for _, fullSources := range []bool{true, false} { + t.Run(map[bool]string{true: "sources", false: "services"}[fullSources], func(t *testing.T) { + l := newCredentialLimiter() + now := time.Now() + l.now = func() time.Time { return now } + key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")} + if fullSources { + ip := netip.MustParseAddr("198.18.0.1") + for range credentialMaxSources { + l.sources[credentialSourceKey{service: key.service, ip: ip}] = &credentialSource{expiresAt: now.Add(credentialBlockDuration), blockedUntil: now.Add(credentialBlockDuration)} + ip = ip.Next() + } + } else { + for i := range credentialMaxServices { + l.services[credentialServiceKey{serviceID: key.service.serviceID, accountID: types.AccountID(strconv.Itoa(i))}] = &credentialService{lastUsed: now} + } + } + _, retry := l.begin(key) + assert.Positive(t, retry, "full state must deny new checks without evicting active entries") + now = now.Add(credentialBlockDuration) + attempt, retry := l.begin(key) + require.Zero(t, retry, "expired state must release capacity") + l.finish(attempt, credentialAccepted) + }) + } +} diff --git a/proxy/internal/auth/credential_test.go b/proxy/internal/auth/credential_test.go new file mode 100644 index 000000000..cae2cff69 --- /dev/null +++ b/proxy/internal/auth/credential_test.go @@ -0,0 +1,196 @@ +package auth + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "net/netip" + "net/url" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "google.golang.org/protobuf/types/known/durationpb" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey" + nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" + "github.com/netbirdio/netbird/management/server/store" + mgmttypes "github.com/netbirdio/netbird/management/server/types" + proxyauth "github.com/netbirdio/netbird/proxy/auth" + "github.com/netbirdio/netbird/proxy/internal/proxy" + "github.com/netbirdio/netbird/shared/management/proto" +) + +// localCredentialClient replaces the transport while keeping the real service +// store, credential verification, and session signing. +type localCredentialClient struct { + server *nbgrpc.ProxyServiceServer +} + +func (c localCredentialClient) Authenticate(ctx context.Context, req *proto.AuthenticateRequest, _ ...grpc.CallOption) (*proto.AuthenticateResponse, error) { + return c.server.Authenticate(ctx, req) +} + +func credentialHandler(t *testing.T, field string) (*Middleware, http.Handler) { + t.Helper() + ctx := context.Background() + s, err := store.NewStore(ctx, mgmttypes.SqliteStoreEngine, t.TempDir(), nil, false) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) }) + require.NoError(t, s.SaveAccount(ctx, &mgmttypes.Account{Id: "account"})) + keys := generateTestKeyPair(t) + svc := &service.Service{ + ID: "service", AccountID: "account", Name: "test", Domain: "example.com", + Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey, + Auth: service.AuthConfig{ + PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"}, + PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "842716"}, + }, + } + require.NoError(t, svc.Auth.HashSecrets()) + require.NoError(t, s.CreateService(ctx, svc)) + server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil) + t.Cleanup(server.Close) + server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil)) + client := localCredentialClient{server: server} + var scheme Scheme = NewPin(client, "service", "account") + if field == "password" { + scheme = NewPassword(client, "service", "account") + } + mw := NewMiddleware(nil, nil, nil) + require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, keys.PublicKey, time.Hour, "account", "service", nil, false, nil)) + return mw, mw.Protect(newPassthroughHandler()) +} + +func credentialRequest(method, field, value string) *http.Request { + r := httptest.NewRequest(method, "https://example.com/", strings.NewReader(url.Values{field: {value}}.Encode())) + r.Header.Set("Content-Type", "application/x-www-form-urlencoded") + r.RemoteAddr = "198.51.100.25:12345" + return r +} + +func TestCredentialAuthPOSTOnly(t *testing.T) { + for _, field := range []string{"pin", "password"} { + t.Run(field, func(t *testing.T) { + _, handler := credentialHandler(t, field) + for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodPatch, http.MethodDelete, http.MethodPost} { + r := credentialRequest(method, field, "") + r.URL.RawQuery = url.Values{field: {"842716"}}.Encode() + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, r) + assert.Equal(t, http.StatusUnauthorized, resp.Code, "%s query credentials must not authenticate", method) + assert.Empty(t, resp.Result().Cookies(), "query credentials must not issue a session") + } + for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodPatch, http.MethodDelete} { + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, credentialRequest(method, field, "842716")) + assert.Equal(t, http.StatusUnauthorized, resp.Code, "%s body credentials must not authenticate", method) + } + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, credentialRequest(http.MethodPost, field, "842716")) + assert.Equal(t, http.StatusSeeOther, resp.Code, "POST body credentials must authenticate") + }) + } +} + +func TestCredentialAuthThrottling(t *testing.T) { + for _, field := range []string{"pin", "password"} { + t.Run(field, func(t *testing.T) { + _, handler := credentialHandler(t, field) + for range 5 { + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, credentialRequest(http.MethodPost, field, "000000")) + require.Equal(t, http.StatusUnauthorized, resp.Code, "initial wrong credentials must be rejected") + } + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, credentialRequest(http.MethodPost, field, "842716")) + assert.Equal(t, http.StatusTooManyRequests, resp.Code, "even correct credentials must wait for the block to expire") + assert.Equal(t, "900", resp.Header().Get("Retry-After"), "five failures must block the source for fifteen minutes") + assert.Empty(t, resp.Result().Cookies(), "blocked credentials must not issue a session") + }) + } +} + +func TestCredentialAuthSessionAndClientIP(t *testing.T) { + keys := generateTestKeyPair(t) + token, err := sessionkey.SignToken(keys.PrivateKey, "pin-user", "", "example.com", proxyauth.MethodPIN, nil, nil, time.Hour) + require.NoError(t, err) + mw := NewMiddleware(nil, nil, nil) + now := time.Now() + mw.credentials.now = func() time.Time { return now } + scheme := &stubScheme{method: proxyauth.MethodPIN, promptID: "pin"} + require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, keys.PublicKey, 0, "account", "service", nil, false, nil)) + handler := mw.Protect(newPassthroughHandler()) + for range credentialFailureLimit { + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, credentialRequest(http.MethodPost, "pin", "000000")) + require.Equal(t, http.StatusUnauthorized, resp.Code, "bad PIN must consume the failure budget") + } + now = now.Add(credentialCheckInterval) + require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, keys.PublicKey, 0, "account", "service", nil, false, nil)) + r := credentialRequest(http.MethodPost, "pin", "000000") + r.RemoteAddr = "[::ffff:198.51.100.25]:45678" + r.Header.Set("X-Forwarded-For", "192.0.2.5") + r.Header.Set("X-Real-IP", "192.0.2.6") + resp := httptest.NewRecorder() + handler.ServeHTTP(resp, r) + assert.Equal(t, http.StatusTooManyRequests, resp.Code, "mapped addresses and untrusted forwarding headers must not bypass the source block") + assert.Equal(t, "no-store", resp.Header().Get("Cache-Control"), "rate limits must not be cached") + r.AddCookie(&http.Cookie{Name: proxyauth.SessionCookieName, Value: token}) + resp = httptest.NewRecorder() + handler.ServeHTTP(resp, r) + assert.Equal(t, http.StatusOK, resp.Code, "an existing session must pass even with credentials in the request") + assert.Equal(t, "backend", resp.Body.String(), "the authenticated request must reach the application") + r = credentialRequest(http.MethodPost, "pin", "000000") + cd := proxy.NewCapturedData("test") + cd.SetClientIP(netip.MustParseAddr("192.0.2.9")) + r = r.WithContext(proxy.WithCapturedData(r.Context(), cd)) + resp = httptest.NewRecorder() + handler.ServeHTTP(resp, r) + assert.Equal(t, http.StatusUnauthorized, resp.Code, "a client resolved by the trusted-proxy middleware must get its own source budget") + r = credentialRequest(http.MethodPost, "pin", "000000") + r.RemoteAddr = "invalid" + resp = httptest.NewRecorder() + handler.ServeHTTP(resp, r) + assert.Equal(t, http.StatusBadRequest, resp.Code, "an unresolvable client address must fail closed") + now = now.Add(credentialBlockDuration) + scheme.token = token + resp = httptest.NewRecorder() + handler.ServeHTTP(resp, credentialRequest(http.MethodPost, "pin", "842716")) + assert.Equal(t, http.StatusSeeOther, resp.Code, "credentials must work again after cooldown") +} + +func TestCredentialAuthManagementThrottling(t *testing.T) { + s, err := status.New(codes.ResourceExhausted, "rate limited").WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(2500 * time.Millisecond)}) + require.NoError(t, err) + for _, tc := range []struct { + name string + err error + code int + retry string + }{ + {"retry info", fmt.Errorf("authenticate PIN: %w", s.Err()), http.StatusTooManyRequests, "3"}, + {"missing retry info", status.Error(codes.ResourceExhausted, "rate limited"), http.StatusTooManyRequests, "6"}, + {"unavailable", status.Error(codes.Unavailable, "unavailable"), http.StatusBadGateway, ""}, + } { + t.Run(tc.name, func(t *testing.T) { + keys := generateTestKeyPair(t) + mw := NewMiddleware(nil, nil, nil) + scheme := &stubScheme{method: proxyauth.MethodPIN, authFn: func(*http.Request) (string, string, error) { return "", "", tc.err }} + require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, keys.PublicKey, 0, "account", "service", nil, false, nil)) + resp := httptest.NewRecorder() + mw.Protect(newPassthroughHandler()).ServeHTTP(resp, credentialRequest(http.MethodPost, "pin", "000000")) + assert.Equal(t, tc.code, resp.Code, "management errors must keep their HTTP meaning") + assert.Equal(t, tc.retry, resp.Header().Get("Retry-After"), "retry hints must round up to whole seconds") + }) + } +} diff --git a/proxy/internal/auth/middleware.go b/proxy/internal/auth/middleware.go index 25ff68010..647741139 100644 --- a/proxy/internal/auth/middleware.go +++ b/proxy/internal/auth/middleware.go @@ -46,7 +46,7 @@ type Scheme interface { // an authenticated user. An empty token indicates an unauthenticated // request; optionally, promptData may be returned for the login UI. // An error indicates an infrastructure failure (e.g. gRPC unavailable). - Authenticate(*http.Request) (token string, promptData string, err error) + Authenticate(*http.Request) (token, promptData string, err error) } // DomainConfig holds the authentication and restriction settings for a protected domain. @@ -77,6 +77,8 @@ type validationResult struct { // Groups for tokens minted before names were embedded; the consumer // falls back to ids for missing positions. GroupNames []string + // MintedToken is the session token issued when a one-time code is redeemed. + MintedToken string } // Middleware applies per-domain authentication and IP restriction checks. @@ -87,6 +89,7 @@ type Middleware struct { sessionValidator SessionValidator geo restrict.GeoResolver tunnelCache *tunnelValidationCache + credentials *credentialLimiter } // NewMiddleware creates a new authentication middleware. The sessionValidator is @@ -101,6 +104,7 @@ func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator, geo re sessionValidator: sessionValidator, geo: geo, tunnelCache: newTunnelValidationCache(), + credentials: newCredentialLimiter(), } } @@ -543,13 +547,9 @@ func (mw *Middleware) authenticateWithSchemes(w http.ResponseWriter, r *http.Req var attemptedMethod string for _, scheme := range config.Schemes { - token, promptData, err := scheme.Authenticate(r) + token, promptData, err := mw.authenticateScheme(r, config, scheme) if err != nil { - mw.logger.WithField("scheme", scheme.Type().String()).Warnf("authentication infrastructure error: %v", err) - if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { - cd.SetOrigin(proxy.OriginAuth) - } - http.Error(w, "authentication service unavailable", http.StatusBadGateway) + mw.writeAuthenticationError(w, r, scheme.Type(), err) return } @@ -583,7 +583,8 @@ func (mw *Middleware) authenticateWithSchemes(w http.ResponseWriter, r *http.Req // handleAuthenticatedToken validates the token, handles denied access, and on // success sets a session cookie and redirects to the original URL. func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Request, host, token string, config DomainConfig, scheme Scheme) { - result, err := mw.validateSessionToken(r.Context(), host, token, config.SessionPublicKey, scheme.Type()) + isCode := scheme.Type() == auth.MethodOIDC && r.URL.Query().Get(auth.SessionCodeQueryParam) != "" + result, err := mw.validateSessionToken(r.Context(), host, token, isCode, config.SessionPublicKey, scheme.Type()) if err != nil { if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { cd.SetOrigin(proxy.OriginAuth) @@ -614,7 +615,13 @@ func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Re return } - setSessionCookie(w, token, config.SessionExpiration) + // When a code was redeemed, the cookie must hold the durable token the + // server returned, not the single-use code. + cookieValue := token + if result.MintedToken != "" { + cookieValue = result.MintedToken + } + setSessionCookie(w, cookieValue, config.SessionExpiration) // Redirect instead of forwarding the auth POST to the backend. // The browser will follow with a GET carrying the new session cookie. @@ -650,11 +657,11 @@ func setSessionCookie(w http.ResponseWriter, token string, expiration time.Durat func wasCredentialSubmitted(r *http.Request, method auth.Method) bool { switch method { case auth.MethodPIN: - return r.FormValue("pin") != "" + return credentialFormValue(r, pinFormId) != "" case auth.MethodPassword: - return r.FormValue("password") != "" + return credentialFormValue(r, passwordFormId) != "" case auth.MethodOIDC: - return r.URL.Query().Get("session_token") != "" + return r.URL.Query().Get(auth.SessionTokenQueryParam) != "" || r.URL.Query().Get(auth.SessionCodeQueryParam) != "" } return false } @@ -708,12 +715,15 @@ func (mw *Middleware) RemoveDomain(domain string) { // validateSessionToken validates a session token. OIDC tokens with a configured // validator go through gRPC for group access checks; other methods validate locally. -func (mw *Middleware) validateSessionToken(ctx context.Context, host, token string, publicKey ed25519.PublicKey, method auth.Method) (*validationResult, error) { +func (mw *Middleware) validateSessionToken(ctx context.Context, host, token string, isCode bool, publicKey ed25519.PublicKey, method auth.Method) (*validationResult, error) { if method == auth.MethodOIDC && mw.sessionValidator != nil { - resp, err := mw.sessionValidator.ValidateSession(ctx, &proto.ValidateSessionRequest{ - Domain: host, - SessionToken: token, - }) + req := &proto.ValidateSessionRequest{Domain: host} + if isCode { + req.SessionCode = token + } else { + req.SessionToken = token //nolint:staticcheck + } + resp, err := mw.sessionValidator.ValidateSession(ctx, req) if err != nil { return nil, fmt.Errorf("%w: %w", errValidationUnavailable, err) } @@ -731,11 +741,12 @@ func (mw *Middleware) validateSessionToken(ctx context.Context, host, token stri }, nil } return &validationResult{ - UserID: resp.UserId, - UserEmail: resp.GetUserEmail(), - Valid: true, - Groups: resp.GetPeerGroupIds(), - GroupNames: resp.GetPeerGroupNames(), + UserID: resp.UserId, + UserEmail: resp.GetUserEmail(), + Valid: true, + Groups: resp.GetPeerGroupIds(), + GroupNames: resp.GetPeerGroupNames(), + MintedToken: resp.GetSessionToken(), }, nil } @@ -790,14 +801,16 @@ func sessionGroupsAllowed(allowed map[string]struct{}, method auth.Method, group } } -// stripSessionTokenParam returns the request URI with the session_token query -// parameter removed so it doesn't linger in the browser's address bar or history. +// stripSessionTokenParam returns the request URI with the session hand-off +// query parameters removed so they don't linger in the browser's address bar +// or history. func stripSessionTokenParam(u *url.URL) string { q := u.Query() - if !q.Has("session_token") { + if !q.Has(auth.SessionTokenQueryParam) && !q.Has(auth.SessionCodeQueryParam) { return u.RequestURI() } - q.Del("session_token") + q.Del(auth.SessionTokenQueryParam) + q.Del(auth.SessionCodeQueryParam) clean := *u clean.RawQuery = q.Encode() return clean.RequestURI() diff --git a/proxy/internal/auth/middleware_test.go b/proxy/internal/auth/middleware_test.go index f1242f95e..cce35ae35 100644 --- a/proxy/internal/auth/middleware_test.go +++ b/proxy/internal/auth/middleware_test.go @@ -783,6 +783,18 @@ func TestWasCredentialSubmitted(t *testing.T) { query: url.Values{"session_token": {"abc123"}}, expected: true, }, + { + name: "OIDC code in query", + method: auth.MethodOIDC, + query: url.Values{"nb_session_code": {"abc123"}}, + expected: true, + }, + { + name: "OIDC backend session_code in query", + method: auth.MethodOIDC, + query: url.Values{"session_code": {"abc123"}}, + expected: false, + }, { name: "OIDC token not in query", method: auth.MethodOIDC, @@ -1571,3 +1583,24 @@ func TestProtect_TunnelPeerFastPath_TakesPathWithInboundMarker(t *testing.T) { assert.Equal(t, http.StatusOK, rec.Code, "a successful tunnel-peer validation must forward to the next handler") } + +func TestStripSessionTokenParam(t *testing.T) { + cases := []struct { + name string + in string + want string + }{ + {"strips session_token", "https://ex.com/p?a=1&session_token=tok", "/p?a=1"}, + {"strips nb_session_code", "https://ex.com/p?a=1&nb_session_code=code", "/p?a=1"}, + {"strips both", "https://ex.com/p?session_token=tok&nb_session_code=code&a=1", "/p?a=1"}, + {"keeps backend session_code", "https://ex.com/p?a=1&session_code=backend", "/p?a=1&session_code=backend"}, + {"no-op when absent", "https://ex.com/p?a=1", "/p?a=1"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + u, err := url.Parse(tc.in) + require.NoError(t, err) + assert.Equal(t, tc.want, stripSessionTokenParam(u)) + }) + } +} diff --git a/proxy/internal/auth/oidc.go b/proxy/internal/auth/oidc.go index a60e6437a..0215fddc3 100644 --- a/proxy/internal/auth/oidc.go +++ b/proxy/internal/auth/oidc.go @@ -40,10 +40,15 @@ func (OIDC) Type() auth.Method { // Authenticate checks for an OIDC session token or obtains the OIDC redirect URL. func (o OIDC) Authenticate(r *http.Request) (string, string, error) { - // Check for the session_token query param (from OIDC redirects). - // The management server passes the token in the URL because it cannot set - // cookies for the proxy's domain (cookies are domain-scoped per RFC 6265). - if token := r.URL.Query().Get("session_token"); token != "" { + // Check for the session credential returned by the OIDC callback. The management + // server passes it in the URL because it cannot set a cookie for the proxy's + // domain (cookies are domain-scoped per RFC 6265). The current flow uses a + // single-use session code to keep the durable token out of the URL. + // session_token remains supported for backward compatibility. + if code := r.URL.Query().Get(auth.SessionCodeQueryParam); code != "" { + return code, "", nil + } + if token := r.URL.Query().Get(auth.SessionTokenQueryParam); token != "" { return token, "", nil } diff --git a/proxy/internal/auth/password.go b/proxy/internal/auth/password.go index 6a7eda3e1..c43e8e3af 100644 --- a/proxy/internal/auth/password.go +++ b/proxy/internal/auth/password.go @@ -35,7 +35,7 @@ func (Password) Type() auth.Method { // so that it can be injected into a request from the UI so that // authentication may be successful. func (p Password) Authenticate(r *http.Request) (string, string, error) { - password := r.FormValue(passwordFormId) + password := credentialFormValue(r, passwordFormId) if password == "" { // No password submitted; return the form ID so the UI can prompt the user. diff --git a/proxy/internal/auth/pin.go b/proxy/internal/auth/pin.go index 4d08f3dc6..180f2648d 100644 --- a/proxy/internal/auth/pin.go +++ b/proxy/internal/auth/pin.go @@ -35,7 +35,7 @@ func (Pin) Type() auth.Method { // so that it can be injected into a request from the UI so that // authentication may be successful. func (p Pin) Authenticate(r *http.Request) (string, string, error) { - pin := r.FormValue(pinFormId) + pin := credentialFormValue(r, pinFormId) if pin == "" { // No PIN submitted; return the form ID so the UI can prompt the user. diff --git a/proxy/internal/metrics/client_metrics_test.go b/proxy/internal/metrics/client_metrics_test.go new file mode 100644 index 000000000..c71e6fb57 --- /dev/null +++ b/proxy/internal/metrics/client_metrics_test.go @@ -0,0 +1,49 @@ +package metrics_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + sdkmetric "go.opentelemetry.io/otel/sdk/metric" + "go.opentelemetry.io/otel/sdk/metric/metricdata" + + "github.com/netbirdio/netbird/proxy/internal/metrics" +) + +func TestRegisterClientObserver(t *testing.T) { + reader := sdkmetric.NewManualReader() + provider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(reader)) + m, err := metrics.New(context.Background(), provider.Meter("test")) + require.NoError(t, err) + + clients := 2 + require.NoError(t, m.RegisterClientObserver(func() int { return clients })) + + var rm metricdata.ResourceMetrics + require.NoError(t, reader.Collect(context.Background(), &rm)) + assert.Equal(t, int64(2), gaugeValue(t, rm, "proxy.clients.count"), "gauge must report the current client count") + + clients = 1 + require.NoError(t, reader.Collect(context.Background(), &rm)) + assert.Equal(t, int64(1), gaugeValue(t, rm, "proxy.clients.count"), "gauge must follow the client count on the next collection") +} + +func gaugeValue(t *testing.T, rm metricdata.ResourceMetrics, name string) int64 { + t.Helper() + + for _, sm := range rm.ScopeMetrics { + for _, mtr := range sm.Metrics { + if mtr.Name != name { + continue + } + gauge, ok := mtr.Data.(metricdata.Gauge[int64]) + require.True(t, ok, "%s must be an int64 gauge", name) + require.Len(t, gauge.DataPoints, 1, "%s must have a single data point", name) + return gauge.DataPoints[0].Value + } + } + t.Fatalf("gauge %s not found", name) + return 0 +} diff --git a/proxy/internal/metrics/metrics.go b/proxy/internal/metrics/metrics.go index 5fd23d934..d7b1797a1 100644 --- a/proxy/internal/metrics/metrics.go +++ b/proxy/internal/metrics/metrics.go @@ -196,6 +196,21 @@ func (m *Metrics) RecordAddPeerDuration(d time.Duration, err error) { )) } +// RegisterClientObserver reports the number of embedded clients as a gauge. +// clientCount runs on every collection cycle, so it must stay cheap. +func (m *Metrics) RegisterClientObserver(clientCount func() int) error { + _, err := m.meter.Int64ObservableGauge( + "proxy.clients.count", + metric.WithUnit("1"), + metric.WithDescription("Current number of embedded NetBird clients running on the netbird proxy"), + metric.WithInt64Callback(func(_ context.Context, o metric.Int64Observer) error { + o.Observe(int64(clientCount())) + return nil + }), + ) + return err +} + func (m *Metrics) initL4Metrics(meter metric.Meter) error { var err error diff --git a/proxy/internal/proxy/reverseproxy.go b/proxy/internal/proxy/reverseproxy.go index 7c9e21261..a3987fe5a 100644 --- a/proxy/internal/proxy/reverseproxy.go +++ b/proxy/internal/proxy/reverseproxy.go @@ -721,12 +721,13 @@ func stripSessionCookie(r *httputil.ProxyRequest) { } } -// stripSessionTokenQuery removes the OIDC session_token query parameter from -// the outgoing URL to prevent credential leakage to backends. +// stripSessionTokenQuery removes the OIDC session hand-off query parameters +// from the outgoing URL to prevent credential leakage to backends. func stripSessionTokenQuery(r *httputil.ProxyRequest) { q := r.Out.URL.Query() - if q.Has("session_token") { - q.Del("session_token") + if q.Has(auth.SessionTokenQueryParam) || q.Has(auth.SessionCodeQueryParam) { + q.Del(auth.SessionTokenQueryParam) + q.Del(auth.SessionCodeQueryParam) r.Out.URL.RawQuery = q.Encode() } } @@ -808,6 +809,12 @@ func classifyProxyError(err error) (title, message string, code int, status web. http.StatusBadGateway, web.ErrorStatus{Proxy: false, Destination: false} + case errors.Is(err, roundtrip.ErrDirectUpstreamBlocked): + return "Destination Not Allowed", + "This proxy does not connect to private or internal addresses. Please contact your administrator.", + http.StatusBadGateway, + web.ErrorStatus{Proxy: false, Destination: false} + case errors.Is(err, roundtrip.ErrTooManyInflight): return "Service Overloaded", "The service is currently handling too many requests. Please try again shortly.", diff --git a/proxy/internal/proxy/reverseproxy_test.go b/proxy/internal/proxy/reverseproxy_test.go index 83afee387..b26ca1f9f 100644 --- a/proxy/internal/proxy/reverseproxy_test.go +++ b/proxy/internal/proxy/reverseproxy_test.go @@ -236,6 +236,17 @@ func TestRewriteFunc_SessionTokenQueryStripping(t *testing.T) { "other query parameters must be preserved") }) + t.Run("strips nb_session_code query parameter", func(t *testing.T) { + pr := newProxyRequest(t, "http://example.com/callback?nb_session_code=code123&other=keep", "1.2.3.4:5000") + + rewrite(pr) + + assert.Empty(t, pr.Out.URL.Query().Get("nb_session_code"), + "OIDC session code must be stripped from backend request") + assert.Equal(t, "keep", pr.Out.URL.Query().Get("other"), + "other query parameters must be preserved") + }) + t.Run("preserves query when no session_token present", func(t *testing.T) { pr := newProxyRequest(t, "http://example.com/api?foo=bar&baz=qux", "1.2.3.4:5000") @@ -1053,6 +1064,17 @@ func TestClassifyProxyError(t *testing.T) { wantCode: http.StatusBadGateway, wantStatus: web.ErrorStatus{Proxy: true, Destination: false}, }, + { + name: "direct upstream blocked by dial guard", + err: &net.OpError{ + Op: "dial", + Net: "tcp", + Err: roundtrip.ErrDirectUpstreamBlocked, + }, + wantTitle: "Destination Not Allowed", + wantCode: http.StatusBadGateway, + wantStatus: web.ErrorStatus{Proxy: false, Destination: false}, + }, { name: "unknown error falls to default", err: errors.New("something unexpected"), diff --git a/proxy/internal/roundtrip/clone_http2_test.go b/proxy/internal/roundtrip/clone_http2_test.go new file mode 100644 index 000000000..f2707ff19 --- /dev/null +++ b/proxy/internal/roundtrip/clone_http2_test.go @@ -0,0 +1,40 @@ +package roundtrip + +import ( + "net/http" + "testing" + + log "github.com/sirupsen/logrus" +) + +// offersHTTP2 reports whether t will negotiate h2 with a TLS upstream. +// Clone forces t's one-time protocol setup, which registers an "h2" +// handler in t.TLSNextProto only when HTTP/2 ended up enabled. +func offersHTTP2(t *http.Transport) bool { + _ = t.Clone() + _, ok := t.TLSNextProto["h2"] + return ok +} + +func TestNewMultiTransportKeepsHTTP2AcrossClone(t *testing.T) { + for _, tc := range []struct { + version string + want bool + }{ + {string(upstreamHTTPAuto), true}, + {string(upstreamHTTP2), true}, + {string(upstreamHTTP11), false}, + } { + t.Run(tc.version, func(t *testing.T) { + t.Setenv(EnvUpstreamHTTPVersion, tc.version) + m := NewMultiTransport(noEmbeddedRoundTripper{}, log.New()) + + if got := offersHTTP2(m.direct.primary); got != tc.want { + t.Errorf("direct transport offers h2 = %v, want %v", got, tc.want) + } + if got := offersHTTP2(m.insecure.primary); got != tc.want { + t.Errorf("insecure transport offers h2 = %v, want %v", got, tc.want) + } + }) + } +} diff --git a/proxy/internal/roundtrip/dialguard.go b/proxy/internal/roundtrip/dialguard.go new file mode 100644 index 000000000..ac01263b4 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard.go @@ -0,0 +1,90 @@ +package roundtrip + +import ( + "context" + "errors" + "net/netip" + "syscall" +) + +// ErrDirectUpstreamBlocked is returned when a direct-upstream dial targets +// an address that is not globally reachable while +// NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE is set. +var ErrDirectUpstreamBlocked = errors.New("direct upstream address is not allowed") + +// blockedUpstreamPrefixes are the ranges that reach the proxy host, its +// cluster or its cloud provider rather than the public internet. NAT64 +// and 6to4 addresses are matched by the IPv4 address they embed. +var blockedUpstreamPrefixes = []netip.Prefix{ + // IPv4 + netip.MustParsePrefix("0.0.0.0/8"), // "this network", including 0.0.0.0 + netip.MustParsePrefix("10.0.0.0/8"), // RFC1918 + netip.MustParsePrefix("100.64.0.0/10"), // CGNAT + netip.MustParsePrefix("127.0.0.0/8"), // loopback + netip.MustParsePrefix("169.254.0.0/16"), // link-local, cloud metadata services + netip.MustParsePrefix("172.16.0.0/12"), // RFC1918 + netip.MustParsePrefix("192.0.0.0/24"), // IETF protocol assignments + netip.MustParsePrefix("192.0.2.0/24"), // documentation + netip.MustParsePrefix("192.88.99.0/24"), // 6to4 relay anycast (deprecated) + netip.MustParsePrefix("192.168.0.0/16"), // RFC1918 + netip.MustParsePrefix("198.18.0.0/15"), // benchmarking + netip.MustParsePrefix("198.51.100.0/24"), // documentation + netip.MustParsePrefix("203.0.113.0/24"), // documentation + netip.MustParsePrefix("224.0.0.0/4"), // multicast + netip.MustParsePrefix("240.0.0.0/4"), // reserved, including broadcast + + // IPv6 + netip.MustParsePrefix("::/96"), // unspecified, loopback, IPv4-compatible + netip.MustParsePrefix("64:ff9b:1::/48"), // local-use NAT64 + netip.MustParsePrefix("100::/64"), // discard-only + netip.MustParsePrefix("2001::/32"), // Teredo + netip.MustParsePrefix("2001:2::/48"), // benchmarking + netip.MustParsePrefix("2001:db8::/32"), // documentation + netip.MustParsePrefix("3fff::/20"), // documentation + netip.MustParsePrefix("5f00::/16"), // SRv6 SIDs + netip.MustParsePrefix("fc00::/7"), // unique local, including AWS IMDS fd00:ec2::254 + netip.MustParsePrefix("fe80::/10"), // link-local + netip.MustParsePrefix("fec0::/10"), // site-local (deprecated) + netip.MustParsePrefix("ff00::/8"), // multicast +} + +var ( + nat64Prefix = netip.MustParsePrefix("64:ff9b::/96") + sixToFour = netip.MustParsePrefix("2002::/16") +) + +// isBlockedUpstreamAddr reports whether a guarded direct-upstream dial +// must refuse addr. +func isBlockedUpstreamAddr(addr netip.Addr) bool { + addr = addr.Unmap().WithZone("") + if !addr.IsValid() { + return true + } + + if nat64Prefix.Contains(addr) { + b := addr.As16() + return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[12:16]))) + } + if sixToFour.Contains(addr) { + b := addr.As16() + return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[2:6]))) + } + + for _, p := range blockedUpstreamPrefixes { + if p.Contains(addr) { + return true + } + } + return false +} + +// guardUpstreamDial is a net.Dialer ControlContext that refuses blocked +// addresses. It sees the resolved address of each socket just before +// connect, so DNS rebinding cannot swap the target after the check. +func guardUpstreamDial(_ context.Context, _, address string, _ syscall.RawConn) error { + ap, err := netip.ParseAddrPort(address) + if err != nil || isBlockedUpstreamAddr(ap.Addr()) { + return ErrDirectUpstreamBlocked + } + return nil +} diff --git a/proxy/internal/roundtrip/dialguard_test.go b/proxy/internal/roundtrip/dialguard_test.go new file mode 100644 index 000000000..79453d8c3 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard_test.go @@ -0,0 +1,189 @@ +package roundtrip + +import ( + "context" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "net/url" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsBlockedUpstreamAddr(t *testing.T) { + blocked := []string{ + "0.0.0.0", + "0.1.2.3", + "10.1.2.3", + "100.64.0.1", + "100.127.255.254", + "127.0.0.1", + "127.255.255.255", + "169.254.169.254", + "172.16.0.1", + "172.31.255.255", + "192.0.0.170", + "192.168.1.1", + "192.88.99.1", + "198.18.0.1", + "224.0.0.1", + "255.255.255.255", + "::", + "::1", + "::169.254.169.254", + "::ffff:127.0.0.1", + "::ffff:169.254.169.254", + "::ffff:10.0.0.1", + "64:ff9b::a9fe:a9fe", // NAT64 of 169.254.169.254 + "64:ff9b::a00:1", // NAT64 of 10.0.0.1 + "64:ff9b:1::1", + "2001::1", + "2001:0:4136:e378:8000:63bf:3fff:fdd2", + "2001:2::1", + "3fff::1", + "5f00::1", + "2002:a9fe:a9fe::1", // 6to4 of 169.254.169.254 + "2002:7f00:1::", // 6to4 of 127.0.0.1 + "fc00::1", + "fd00:ec2::254", + "fe80::1", + "fe80::1%eth0", + "fec0::1", + "ff02::1", + } + for _, s := range blocked { + t.Run("blocks "+s, func(t *testing.T) { + assert.True(t, isBlockedUpstreamAddr(netip.MustParseAddr(s))) + }) + } + + allowed := []string{ + "1.1.1.1", + "8.8.8.8", + "100.63.255.255", + "100.128.0.0", + "172.15.255.255", + "172.32.0.0", + "169.253.255.255", + "2606:4700:4700::1111", + "2001:4860:4860::8888", + "2001:1::1", + "4000::1", + "::ffff:8.8.8.8", + "64:ff9b::808:808", // NAT64 of 8.8.8.8 + "2002:808:808::1", // 6to4 of 8.8.8.8 + } + for _, s := range allowed { + t.Run("allows "+s, func(t *testing.T) { + assert.False(t, isBlockedUpstreamAddr(netip.MustParseAddr(s))) + }) + } + + assert.True(t, isBlockedUpstreamAddr(netip.Addr{}), "the zero Addr must be refused") +} + +func TestGuardUpstreamDial_RejectsUnparsableAddress(t *testing.T) { + err := guardUpstreamDial(context.Background(), "tcp", "not-an-address", nil) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "an address the guard cannot parse must fail closed") +} + +// TestMultiTransport_BlockPrivateUpstreams exercises the guard end to end +// against a loopback test server: by IP literal and by a hostname that +// resolves to loopback, on both direct branches, and confirms the +// embedded branch is not affected. +func TestMultiTransport_BlockPrivateUpstreams(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, "reached") + })) + defer srv.Close() + + _, port, err := net.SplitHostPort(srv.Listener.Addr().String()) + require.NoError(t, err) + byName := (&url.URL{Scheme: "http", Host: net.JoinHostPort("localhost", port)}).String() + + directCtx := WithDirectUpstream(context.Background()) + insecureCtx := WithSkipTLSVerify(directCtx) + + // roundTrip returns the response body, so callers never hold one open. + roundTrip := func(t *testing.T, mt *MultiTransport, ctx context.Context, target string) (string, error) { + t.Helper() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil) + require.NoError(t, err) + resp, err := mt.RoundTrip(req) + if err != nil { + return "", err + } + defer func() { _ = resp.Body.Close() }() + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return string(body), nil + } + + t.Run("enabled refuses loopback", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "true") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + cases := []struct { + name string + ctx context.Context + target string + }{ + {"direct by IP", directCtx, srv.URL}, + {"direct by hostname", directCtx, byName}, + {"insecure by IP", insecureCtx, srv.URL}, + {"insecure by hostname", insecureCtx, byName}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := roundTrip(t, mt, tc.ctx, tc.target) + require.Error(t, err) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked) + }) + } + }) + + t.Run("invalid value enables the guard", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "yes please") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + _, err := roundTrip(t, mt, directCtx, srv.URL) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "a value that does not parse must fail closed") + }) + + t.Run("explicit false disables the guard", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "false") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + body, err := roundTrip(t, mt, directCtx, srv.URL) + require.NoError(t, err) + assert.Equal(t, "reached", body) + }) + + t.Run("enabled leaves embedded branch alone", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "true") + embedded := &stubRoundTripper{body: "embedded"} + mt := NewMultiTransport(embedded, nil) + + body, err := roundTrip(t, mt, context.Background(), srv.URL) + require.NoError(t, err) + assert.Equal(t, "embedded", body) + assert.True(t, embedded.called, "the guard must not change dispatch to the embedded transport") + }) + + t.Run("disabled by default", func(t *testing.T) { + // Register the restore first so an exported value comes back after + // the test, then exercise a genuinely absent variable. + t.Setenv(EnvDirectUpstreamBlockPrivate, "") + require.NoError(t, os.Unsetenv(EnvDirectUpstreamBlockPrivate)) + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + body, err := roundTrip(t, mt, directCtx, srv.URL) + require.NoError(t, err, "private and self-hosted proxies must keep reaching local upstreams") + assert.Equal(t, "reached", body) + }) +} diff --git a/proxy/internal/roundtrip/multi.go b/proxy/internal/roundtrip/multi.go index 1abf54a8d..a430d45bd 100644 --- a/proxy/internal/roundtrip/multi.go +++ b/proxy/internal/roundtrip/multi.go @@ -41,7 +41,9 @@ var errNoEmbeddedTransport = errors.New("multitransport: embedded roundtripper n // MultiTransport that only ever uses the direct branch. The direct // branches honour the same NB_PROXY_* tuning env vars as the embedded // transport (see loadTransportConfig) plus a dial-timeout wrapper that -// respects types.WithDialTimeout. +// respects types.WithDialTimeout. With NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE +// set, the direct branches refuse addresses that are not globally reachable +// (see guardUpstreamDial). func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTransport { if logger == nil { logger = log.StandardLogger() @@ -51,6 +53,9 @@ func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTra Timeout: 30 * time.Second, KeepAlive: 30 * time.Second, } + if cfg.blockPrivateUpstreams { + dialer.ControlContext = guardUpstreamDial + } direct := &http.Transport{ DialContext: dialWithTimeout(dialer.DialContext), MaxIdleConns: cfg.maxIdleConns, @@ -64,6 +69,9 @@ func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTra ReadBufferSize: cfg.readBufferSize, DisableCompression: cfg.disableCompression, } + // Clone runs the transport's one-time protocol setup, so the HTTP + // version must be applied first or the source loses HTTP/2 for good. + applyUpstreamHTTPVersion(direct, cfg.upstreamHTTPVersion) insecure := direct.Clone() insecure.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} //nolint:gosec // matches the embedded NetBird transport's per-target opt-in diff --git a/proxy/internal/roundtrip/netbird.go b/proxy/internal/roundtrip/netbird.go index d7b464182..07b497c46 100644 --- a/proxy/internal/roundtrip/netbird.go +++ b/proxy/internal/roundtrip/netbird.go @@ -425,6 +425,9 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account ReadBufferSize: n.transportCfg.readBufferSize, DisableCompression: n.transportCfg.disableCompression, } + // Clone runs the transport's one-time protocol setup, so the HTTP + // version must be applied first or the source loses HTTP/2 for good. + applyUpstreamHTTPVersion(transport, n.transportCfg.upstreamHTTPVersion) insecureTransport := transport.Clone() insecureTransport.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} //nolint:gosec diff --git a/proxy/internal/roundtrip/transport.go b/proxy/internal/roundtrip/transport.go index 9e872e447..6383079c4 100644 --- a/proxy/internal/roundtrip/transport.go +++ b/proxy/internal/roundtrip/transport.go @@ -25,6 +25,12 @@ const ( EnvDisableCompression = "NB_PROXY_DISABLE_COMPRESSION" EnvMaxInflight = "NB_PROXY_MAX_INFLIGHT" EnvUpstreamHTTPVersion = "NB_PROXY_UPSTREAM_HTTP_VERSION" + // EnvDirectUpstreamBlockPrivate refuses direct-upstream dials to + // addresses that are not globally reachable (loopback, private, + // link-local, CGNAT, ...). Off by default: private and self-hosted + // proxies use direct_upstream to reach LAN and localhost services. + // Proxies that serve untrusted accounts must turn it on. + EnvDirectUpstreamBlockPrivate = "NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE" ) // upstreamHTTPVersion selects the HTTP version the proxy uses towards an @@ -69,6 +75,9 @@ type transportConfig struct { // explicit values are for backends whose advertised h2 support is // unusable and whose failure mode the negotiation cannot see. upstreamHTTPVersion upstreamHTTPVersion + // blockPrivateUpstreams guards the direct branches' dialer with + // guardUpstreamDial. It has no effect on the embedded branch. + blockPrivateUpstreams bool } func defaultTransportConfig() transportConfig { @@ -122,6 +131,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig { if v, ok := envUpstreamHTTPVersion(EnvUpstreamHTTPVersion, logger); ok { cfg.upstreamHTTPVersion = v } + cfg.blockPrivateUpstreams = envGuardBool(EnvDirectUpstreamBlockPrivate, logger) logger.WithFields(log.Fields{ "max_idle_conns": cfg.maxIdleConns, @@ -136,6 +146,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig { "disable_compression": cfg.disableCompression, "max_inflight": cfg.maxInflight, "upstream_http_version": cfg.upstreamHTTPVersion, + "block_private_upstreams": cfg.blockPrivateUpstreams, }).Debug("backend transport configuration") return cfg @@ -246,6 +257,22 @@ func envDuration(key string, logger *log.Logger) (time.Duration, bool) { return v, true } +// envGuardBool reads a bool that turns a security guard on. Unset means +// off, but a value that does not parse turns the guard on: a typo must not +// leave a proxy that was meant to be guarded without the guard. +func envGuardBool(key string, logger *log.Logger) bool { + s := os.Getenv(key) + if s == "" { + return false + } + v, err := strconv.ParseBool(s) + if err != nil { + logger.Warnf("failed to parse %s=%q as bool, enabling it: %v", key, s, err) + return true + } + return v +} + func envBool(key string, logger *log.Logger) (bool, bool) { s := os.Getenv(key) if s == "" { diff --git a/proxy/management_byop_integration_test.go b/proxy/management_byop_integration_test.go index d075e47ec..42301254d 100644 --- a/proxy/management_byop_integration_test.go +++ b/proxy/management_byop_integration_test.go @@ -104,7 +104,7 @@ func setupBYOPIntegrationTest(t *testing.T) *byopTestSetup { require.NoError(t, err) tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) - pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore) meter := noop.NewMeterProvider().Meter("test") realProxyManager, err := proxymanager.NewManager(testStore, meter) @@ -121,7 +121,7 @@ func setupBYOPIntegrationTest(t *testing.T) *byopTestSetup { proxyService := nbgrpc.NewProxyServiceServer( &testAccessLogManager{}, tokenStore, - pkceStore, + singleUseStore, oidcConfig, nil, usersManager, diff --git a/proxy/management_integration_test.go b/proxy/management_integration_test.go index df016e790..03a9855de 100644 --- a/proxy/management_integration_test.go +++ b/proxy/management_integration_test.go @@ -119,7 +119,7 @@ func setupIntegrationTest(t *testing.T) *integrationTestSetup { require.NoError(t, err) tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) - pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) + singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore) // Create real users manager usersManager := users.NewManager(testStore) @@ -131,12 +131,12 @@ func setupIntegrationTest(t *testing.T) *integrationTestSetup { HMACKey: []byte("test-hmac-key"), } - proxyManager := &testProxyManager{} + proxyManager := &testProxyManager{supportsSessionCode: true} proxyService := nbgrpc.NewProxyServiceServer( &testAccessLogManager{}, tokenStore, - pkceStore, + singleUseStore, oidcConfig, nil, usersManager, @@ -202,9 +202,11 @@ func (m *testAccessLogManager) GetAllAccessLogs(_ context.Context, _, _ string, } // testProxyManager is a mock implementation of proxy.Manager for testing. -type testProxyManager struct{} +type testProxyManager struct { + supportsSessionCode bool +} -func (m *testProxyManager) Connect(_ context.Context, proxyID, sessionID, _, _ string, _ *string, _ *nbproxy.Capabilities) (*nbproxy.Proxy, error) { +func (m *testProxyManager) Connect(_ context.Context, proxyID, sessionID, _, _, _ string, _ *string, _ *nbproxy.Capabilities) (*nbproxy.Proxy, error) { return &nbproxy.Proxy{ID: proxyID, SessionID: sessionID, Status: nbproxy.StatusConnected}, nil } @@ -244,6 +246,14 @@ func (m *testProxyManager) ClusterSupportsPrivate(_ context.Context, _ string) * return nil } +func (m *testProxyManager) ClusterAllProxiesPrivate(_ context.Context, _ string) *bool { + return nil +} + +func (m *testProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool { + return m.supportsSessionCode +} + func (m *testProxyManager) CleanupStale(_ context.Context, _ time.Duration) error { return nil } diff --git a/proxy/server.go b/proxy/server.go index 5b652e61c..762ead9b8 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -362,6 +362,13 @@ func (s *Server) Start(ctx context.Context) error { return err } + startupOK := false + defer func() { + if !startupOK { + s.cleanupFailedStart() + } + }() + // Management client must be initialised BEFORE the middleware manager — // initMiddlewareManager passes s.mgmtClient into the builtin FactoryContext // that the limit-check / limit-record middlewares pull from. Reversed @@ -374,7 +381,9 @@ func (s *Server) Start(ctx context.Context) error { runCtx, runCancel := context.WithCancel(ctx) s.runCancel = runCancel - s.initNetBirdClient() + if err := s.initNetBirdClient(); err != nil { + return err + } // Create health checker before the mapping worker so it can track // management connectivity from the first stream connection. s.healthChecker = health.NewChecker(s.Logger, s.netbird) @@ -395,18 +404,6 @@ func (s *Server) Start(ctx context.Context) error { return err } - startupOK := false - defer func() { - if startupOK { - return - } - if s.geoRaw != nil { - if closeErr := s.geoRaw.Close(); closeErr != nil { - s.Logger.Debugf("close geolocation on startup failure: %v", closeErr) - } - } - }() - s.auth = auth.NewMiddleware(s.Logger, s.mgmtClient, s.geo) s.accessLog = accesslog.NewLogger(s.mgmtClient, s.Logger, s.TrustedProxies) @@ -475,14 +472,7 @@ func (s *Server) Stop(ctx context.Context) error { go func() { defer close(done) s.gracefulShutdown() - if s.runCancel != nil { - s.runCancel() - } - if s.mgmtConn != nil { - if err := s.mgmtConn.Close(); err != nil { - s.Logger.Debugf("management connection close: %v", err) - } - } + s.releaseRunResources() }() select { @@ -497,6 +487,27 @@ func (s *Server) Stop(ctx context.Context) error { return s.runErr } +// cleanupFailedStart releases what a failed Start already brought up. It +// skips the drain and pre-stop delay because nothing has served yet, and +// consumes stopOnce so a later Stop stays a no-op. +func (s *Server) cleanupFailedStart() { + s.stopOnce.Do(func() { + s.shutdownServices() + s.releaseRunResources() + }) +} + +func (s *Server) releaseRunResources() { + if s.runCancel != nil { + s.runCancel() + } + if s.mgmtConn != nil { + if err := s.mgmtConn.Close(); err != nil { + s.Logger.Debugf("management connection close: %v", err) + } + } +} + // waitAndStop blocks until ctx is cancelled or a background goroutine // reports a fatal error, then drains and stops. Used by ListenAndServe. func (s *Server) waitAndStop(ctx context.Context) error { @@ -568,7 +579,7 @@ func (s *Server) initManagementClient() error { // initNetBirdClient builds the multi-tenant embedded NetBird client used // for outbound RoundTripping and (when --private is on) per-account // inbound listeners. -func (s *Server) initNetBirdClient() { +func (s *Server) initNetBirdClient() error { s.netbird = roundtrip.NewNetBird(s.ctx, s.ID, s.ProxyURL, roundtrip.ClientConfig{ MgmtAddr: s.ManagementAddress, WGPort: s.WireguardPort, @@ -581,6 +592,10 @@ func (s *Server) initNetBirdClient() { BlockInbound: !s.Private, }, s.Logger, s, s.mgmtClient) s.netbird.OnAddPeer = s.meter.RecordAddPeerDuration + if err := s.meter.RegisterClientObserver(s.netbird.ClientCount); err != nil { + return fmt.Errorf("register client metrics: %w", err) + } + return nil } // initReverseProxy builds the meter-instrumented reverse proxy. MultiTransport diff --git a/proxy/server_test.go b/proxy/server_test.go index 9cef63b95..cf583985f 100644 --- a/proxy/server_test.go +++ b/proxy/server_test.go @@ -16,6 +16,7 @@ import ( "github.com/stretchr/testify/require" "go.opentelemetry.io/otel/metric/noop" "google.golang.org/grpc" + "google.golang.org/grpc/connectivity" "github.com/netbirdio/netbird/proxy/internal/auth" proxymetrics "github.com/netbirdio/netbird/proxy/internal/metrics" @@ -106,6 +107,25 @@ func TestStartFailsWithoutManagement(t *testing.T) { assert.Contains(t, err.Error(), "already started", "error must explain why the call was rejected") } +func TestStartFailureReleasesManagementConnection(t *testing.T) { + srv := New(t.Context(), Config{ + Logger: quietLifecycleLogger(), + ListenAddr: "127.0.0.1:0", + ManagementAddress: "https://127.0.0.1:1", + CertificateDirectory: t.TempDir(), + CertificateFile: "missing.crt", + CertificateKeyFile: "missing.key", + }) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + err := srv.Start(ctx) + require.Error(t, err, "Start must fail on the missing certificate") + require.NotNil(t, srv.mgmtConn, "the management connection is created before the certificate step") + assert.Equal(t, connectivity.Shutdown, srv.mgmtConn.GetState(), "a failed Start must close the management connection it opened") +} + func TestStopIsIdempotent(t *testing.T) { srv := &Server{ Logger: quietLifecycleLogger(), diff --git a/proxy/web/dist/assets/index.js b/proxy/web/dist/assets/index.js index 9ce3e4394..0a34a21d4 100644 --- a/proxy/web/dist/assets/index.js +++ b/proxy/web/dist/assets/index.js @@ -1,9 +1,9 @@ -(function(){const v=document.createElement("link").relList;if(v&&v.supports&&v.supports("modulepreload"))return;for(const _ of document.querySelectorAll('link[rel="modulepreload"]'))f(_);new MutationObserver(_=>{for(const O of _)if(O.type==="childList")for(const D of O.addedNodes)D.tagName==="LINK"&&D.rel==="modulepreload"&&f(D)}).observe(document,{childList:!0,subtree:!0});function S(_){const O={};return _.integrity&&(O.integrity=_.integrity),_.referrerPolicy&&(O.referrerPolicy=_.referrerPolicy),_.crossOrigin==="use-credentials"?O.credentials="include":_.crossOrigin==="anonymous"?O.credentials="omit":O.credentials="same-origin",O}function f(_){if(_.ep)return;_.ep=!0;const O=S(_);fetch(_.href,O)}})();var Sf={exports:{}},Du={};var Yd;function jm(){if(Yd)return Du;Yd=1;var r=Symbol.for("react.transitional.element"),v=Symbol.for("react.fragment");function S(f,_,O){var D=null;if(O!==void 0&&(D=""+O),_.key!==void 0&&(D=""+_.key),"key"in _){O={};for(var U in _)U!=="key"&&(O[U]=_[U])}else O=_;return _=O.ref,{$$typeof:r,type:f,key:D,ref:_!==void 0?_:null,props:O}}return Du.Fragment=v,Du.jsx=S,Du.jsxs=S,Du}var Gd;function Rm(){return Gd||(Gd=1,Sf.exports=jm()),Sf.exports}var A=Rm(),xf={exports:{}},K={};var Xd;function Hm(){if(Xd)return K;Xd=1;var r=Symbol.for("react.transitional.element"),v=Symbol.for("react.portal"),S=Symbol.for("react.fragment"),f=Symbol.for("react.strict_mode"),_=Symbol.for("react.profiler"),O=Symbol.for("react.consumer"),D=Symbol.for("react.context"),U=Symbol.for("react.forward_ref"),N=Symbol.for("react.suspense"),p=Symbol.for("react.memo"),R=Symbol.for("react.lazy"),H=Symbol.for("react.activity"),V=Symbol.iterator;function st(s){return s===null||typeof s!="object"?null:(s=V&&s[V]||s["@@iterator"],typeof s=="function"?s:null)}var ct={isMounted:function(){return!1},enqueueForceUpdate:function(){},enqueueReplaceState:function(){},enqueueSetState:function(){}},G=Object.assign,Q={};function L(s,M,j){this.props=s,this.context=M,this.refs=Q,this.updater=j||ct}L.prototype.isReactComponent={},L.prototype.setState=function(s,M){if(typeof s!="object"&&typeof s!="function"&&s!=null)throw Error("takes an object of state variables to update or a function which returns an object of state variables.");this.updater.enqueueSetState(this,s,M,"setState")},L.prototype.forceUpdate=function(s){this.updater.enqueueForceUpdate(this,s,"forceUpdate")};function gt(){}gt.prototype=L.prototype;function zt(s,M,j){this.props=s,this.context=M,this.refs=Q,this.updater=j||ct}var _t=zt.prototype=new gt;_t.constructor=zt,G(_t,L.prototype),_t.isPureReactComponent=!0;var it=Array.isArray;function Ot(){}var J={H:null,A:null,T:null,S:null},Rt=Object.prototype.hasOwnProperty;function It(s,M,j){var q=j.ref;return{$$typeof:r,type:s,key:M,ref:q!==void 0?q:null,props:j}}function jl(s,M){return It(s.type,M,s.props)}function Pt(s){return typeof s=="object"&&s!==null&&s.$$typeof===r}function I(s){var M={"=":"=0",":":"=2"};return"$"+s.replace(/[=:]/g,function(j){return M[j]})}var Rl=/\/+/g;function tl(s,M){return typeof s=="object"&&s!==null&&s.key!=null?I(""+s.key):M.toString(36)}function ll(s){switch(s.status){case"fulfilled":return s.value;case"rejected":throw s.reason;default:switch(typeof s.status=="string"?s.then(Ot,Ot):(s.status="pending",s.then(function(M){s.status==="pending"&&(s.status="fulfilled",s.value=M)},function(M){s.status==="pending"&&(s.status="rejected",s.reason=M)})),s.status){case"fulfilled":return s.value;case"rejected":throw s.reason}}throw s}function x(s,M,j,q,k){var P=typeof s;(P==="undefined"||P==="boolean")&&(s=null);var yt=!1;if(s===null)yt=!0;else switch(P){case"bigint":case"string":case"number":yt=!0;break;case"object":switch(s.$$typeof){case r:case v:yt=!0;break;case R:return yt=s._init,x(yt(s._payload),M,j,q,k)}}if(yt)return k=k(s),yt=q===""?"."+tl(s,0):q,it(k)?(j="",yt!=null&&(j=yt.replace(Rl,"$&/")+"/"),x(k,M,j,"",function(qa){return qa})):k!=null&&(Pt(k)&&(k=jl(k,j+(k.key==null||s&&s.key===k.key?"":(""+k.key).replace(Rl,"$&/")+"/")+yt)),M.push(k)),1;yt=0;var Wt=q===""?".":q+":";if(it(s))for(var Ut=0;Ut>>1,dt=x[nt];if(0<_(dt,C))x[nt]=C,x[Z]=dt,Z=nt;else break t}}function S(x){return x.length===0?null:x[0]}function f(x){if(x.length===0)return null;var C=x[0],Z=x.pop();if(Z!==C){x[0]=Z;t:for(var nt=0,dt=x.length,s=dt>>>1;nt_(j,Z))q_(k,j)?(x[nt]=k,x[q]=Z,nt=q):(x[nt]=j,x[M]=Z,nt=M);else if(q_(k,Z))x[nt]=k,x[q]=Z,nt=q;else break t}}return C}function _(x,C){var Z=x.sortIndex-C.sortIndex;return Z!==0?Z:x.id-C.id}if(r.unstable_now=void 0,typeof performance=="object"&&typeof performance.now=="function"){var O=performance;r.unstable_now=function(){return O.now()}}else{var D=Date,U=D.now();r.unstable_now=function(){return D.now()-U}}var N=[],p=[],R=1,H=null,V=3,st=!1,ct=!1,G=!1,Q=!1,L=typeof setTimeout=="function"?setTimeout:null,gt=typeof clearTimeout=="function"?clearTimeout:null,zt=typeof setImmediate<"u"?setImmediate:null;function _t(x){for(var C=S(p);C!==null;){if(C.callback===null)f(p);else if(C.startTime<=x)f(p),C.sortIndex=C.expirationTime,v(N,C);else break;C=S(p)}}function it(x){if(G=!1,_t(x),!ct)if(S(N)!==null)ct=!0,Ot||(Ot=!0,I());else{var C=S(p);C!==null&&ll(it,C.startTime-x)}}var Ot=!1,J=-1,Rt=5,It=-1;function jl(){return Q?!0:!(r.unstable_now()-Itx&&jl());){var nt=H.callback;if(typeof nt=="function"){H.callback=null,V=H.priorityLevel;var dt=nt(H.expirationTime<=x);if(x=r.unstable_now(),typeof dt=="function"){H.callback=dt,_t(x),C=!0;break l}H===S(N)&&f(N),_t(x)}else f(N);H=S(N)}if(H!==null)C=!0;else{var s=S(p);s!==null&&ll(it,s.startTime-x),C=!1}}break t}finally{H=null,V=Z,st=!1}C=void 0}}finally{C?I():Ot=!1}}}var I;if(typeof zt=="function")I=function(){zt(Pt)};else if(typeof MessageChannel<"u"){var Rl=new MessageChannel,tl=Rl.port2;Rl.port1.onmessage=Pt,I=function(){tl.postMessage(null)}}else I=function(){L(Pt,0)};function ll(x,C){J=L(function(){x(r.unstable_now())},C)}r.unstable_IdlePriority=5,r.unstable_ImmediatePriority=1,r.unstable_LowPriority=4,r.unstable_NormalPriority=3,r.unstable_Profiling=null,r.unstable_UserBlockingPriority=2,r.unstable_cancelCallback=function(x){x.callback=null},r.unstable_forceFrameRate=function(x){0>x||125nt?(x.sortIndex=Z,v(p,x),S(N)===null&&x===S(p)&&(G?(gt(J),J=-1):G=!0,ll(it,Z-nt))):(x.sortIndex=dt,v(N,x),ct||st||(ct=!0,Ot||(Ot=!0,I()))),x},r.unstable_shouldYield=jl,r.unstable_wrapCallback=function(x){var C=V;return function(){var Z=V;V=C;try{return x.apply(this,arguments)}finally{V=Z}}}})(Ef)),Ef}var wd;function qm(){return wd||(wd=1,Tf.exports=Bm()),Tf.exports}var Af={exports:{}},kt={};var Ld;function Ym(){if(Ld)return kt;Ld=1;var r=Rf();function v(N){var p="https://react.dev/errors/"+N;if(1"u"||typeof __REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE!="function"))try{__REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE(r)}catch(v){console.error(v)}}return r(),Af.exports=Ym(),Af.exports}var Kd;function Xm(){if(Kd)return Uu;Kd=1;var r=qm(),v=Rf(),S=Gm();function f(t){var l="https://react.dev/errors/"+t;if(1dt||(t.current=nt[dt],nt[dt]=null,dt--)}function j(t,l){dt++,nt[dt]=t.current,t.current=l}var q=s(null),k=s(null),P=s(null),yt=s(null);function Wt(t,l){switch(j(P,l),j(k,t),j(q,null),l.nodeType){case 9:case 11:t=(t=l.documentElement)&&(t=t.namespaceURI)?cd(t):0;break;default:if(t=l.tagName,l=l.namespaceURI)l=cd(l),t=fd(l,t);else switch(t){case"svg":t=1;break;case"math":t=2;break;default:t=0}}M(q),j(q,t)}function Ut(){M(q),M(k),M(P)}function qa(t){t.memoizedState!==null&&j(yt,t);var l=q.current,e=fd(l,t.type);l!==e&&(j(k,t),j(q,e))}function Hu(t){k.current===t&&(M(q),M(k)),yt.current===t&&(M(yt),Mu._currentValue=Z)}var li,Bf;function Ue(t){if(li===void 0)try{throw Error()}catch(e){var l=e.stack.trim().match(/\n( *(at )?)/);li=l&&l[1]||"",Bf=-1{for(const O of _)if(O.type==="childList")for(const D of O.addedNodes)D.tagName==="LINK"&&D.rel==="modulepreload"&&f(D)}).observe(document,{childList:!0,subtree:!0});function S(_){const O={};return _.integrity&&(O.integrity=_.integrity),_.referrerPolicy&&(O.referrerPolicy=_.referrerPolicy),_.crossOrigin==="use-credentials"?O.credentials="include":_.crossOrigin==="anonymous"?O.credentials="omit":O.credentials="same-origin",O}function f(_){if(_.ep)return;_.ep=!0;const O=S(_);fetch(_.href,O)}})();var Sf={exports:{}},Du={};var Yd;function jm(){if(Yd)return Du;Yd=1;var r=Symbol.for("react.transitional.element"),v=Symbol.for("react.fragment");function S(f,_,O){var D=null;if(O!==void 0&&(D=""+O),_.key!==void 0&&(D=""+_.key),"key"in _){O={};for(var U in _)U!=="key"&&(O[U]=_[U])}else O=_;return _=O.ref,{$$typeof:r,type:f,key:D,ref:_!==void 0?_:null,props:O}}return Du.Fragment=v,Du.jsx=S,Du.jsxs=S,Du}var Gd;function Rm(){return Gd||(Gd=1,Sf.exports=jm()),Sf.exports}var E=Rm(),xf={exports:{}},K={};var Xd;function Hm(){if(Xd)return K;Xd=1;var r=Symbol.for("react.transitional.element"),v=Symbol.for("react.portal"),S=Symbol.for("react.fragment"),f=Symbol.for("react.strict_mode"),_=Symbol.for("react.profiler"),O=Symbol.for("react.consumer"),D=Symbol.for("react.context"),U=Symbol.for("react.forward_ref"),N=Symbol.for("react.suspense"),p=Symbol.for("react.memo"),R=Symbol.for("react.lazy"),H=Symbol.for("react.activity"),L=Symbol.iterator;function ot(o){return o===null||typeof o!="object"?null:(o=L&&o[L]||o["@@iterator"],typeof o=="function"?o:null)}var ct={isMounted:function(){return!1},enqueueForceUpdate:function(){},enqueueReplaceState:function(){},enqueueSetState:function(){}},G=Object.assign,Q={};function V(o,M,j){this.props=o,this.context=M,this.refs=Q,this.updater=j||ct}V.prototype.isReactComponent={},V.prototype.setState=function(o,M){if(typeof o!="object"&&typeof o!="function"&&o!=null)throw Error("takes an object of state variables to update or a function which returns an object of state variables.");this.updater.enqueueSetState(this,o,M,"setState")},V.prototype.forceUpdate=function(o){this.updater.enqueueForceUpdate(this,o,"forceUpdate")};function gt(){}gt.prototype=V.prototype;function zt(o,M,j){this.props=o,this.context=M,this.refs=Q,this.updater=j||ct}var _t=zt.prototype=new gt;_t.constructor=zt,G(_t,V.prototype),_t.isPureReactComponent=!0;var nt=Array.isArray;function Ot(){}var J={H:null,A:null,T:null,S:null},Nt=Object.prototype.hasOwnProperty;function Xt(o,M,j){var q=j.ref;return{$$typeof:r,type:o,key:M,ref:q!==void 0?q:null,props:j}}function pl(o,M){return Xt(o.type,M,o.props)}function Pt(o){return typeof o=="object"&&o!==null&&o.$$typeof===r}function I(o){var M={"=":"=0",":":"=2"};return"$"+o.replace(/[=:]/g,function(j){return M[j]})}var Rl=/\/+/g;function tl(o,M){return typeof o=="object"&&o!==null&&o.key!=null?I(""+o.key):M.toString(36)}function ll(o){switch(o.status){case"fulfilled":return o.value;case"rejected":throw o.reason;default:switch(typeof o.status=="string"?o.then(Ot,Ot):(o.status="pending",o.then(function(M){o.status==="pending"&&(o.status="fulfilled",o.value=M)},function(M){o.status==="pending"&&(o.status="rejected",o.reason=M)})),o.status){case"fulfilled":return o.value;case"rejected":throw o.reason}}throw o}function x(o,M,j,q,k){var P=typeof o;(P==="undefined"||P==="boolean")&&(o=null);var yt=!1;if(o===null)yt=!0;else switch(P){case"bigint":case"string":case"number":yt=!0;break;case"object":switch(o.$$typeof){case r:case v:yt=!0;break;case R:return yt=o._init,x(yt(o._payload),M,j,q,k)}}if(yt)return k=k(o),yt=q===""?"."+tl(o,0):q,nt(k)?(j="",yt!=null&&(j=yt.replace(Rl,"$&/")+"/"),x(k,M,j,"",function(qa){return qa})):k!=null&&(Pt(k)&&(k=pl(k,j+(k.key==null||o&&o.key===k.key?"":(""+k.key).replace(Rl,"$&/")+"/")+yt)),M.push(k)),1;yt=0;var $t=q===""?".":q+":";if(nt(o))for(var Ct=0;Ct>>1,dt=x[it];if(0<_(dt,C))x[it]=C,x[Z]=dt,Z=it;else break t}}function S(x){return x.length===0?null:x[0]}function f(x){if(x.length===0)return null;var C=x[0],Z=x.pop();if(Z!==C){x[0]=Z;t:for(var it=0,dt=x.length,o=dt>>>1;it_(j,Z))q_(k,j)?(x[it]=k,x[q]=Z,it=q):(x[it]=j,x[M]=Z,it=M);else if(q_(k,Z))x[it]=k,x[q]=Z,it=q;else break t}}return C}function _(x,C){var Z=x.sortIndex-C.sortIndex;return Z!==0?Z:x.id-C.id}if(r.unstable_now=void 0,typeof performance=="object"&&typeof performance.now=="function"){var O=performance;r.unstable_now=function(){return O.now()}}else{var D=Date,U=D.now();r.unstable_now=function(){return D.now()-U}}var N=[],p=[],R=1,H=null,L=3,ot=!1,ct=!1,G=!1,Q=!1,V=typeof setTimeout=="function"?setTimeout:null,gt=typeof clearTimeout=="function"?clearTimeout:null,zt=typeof setImmediate<"u"?setImmediate:null;function _t(x){for(var C=S(p);C!==null;){if(C.callback===null)f(p);else if(C.startTime<=x)f(p),C.sortIndex=C.expirationTime,v(N,C);else break;C=S(p)}}function nt(x){if(G=!1,_t(x),!ct)if(S(N)!==null)ct=!0,Ot||(Ot=!0,I());else{var C=S(p);C!==null&&ll(nt,C.startTime-x)}}var Ot=!1,J=-1,Nt=5,Xt=-1;function pl(){return Q?!0:!(r.unstable_now()-Xtx&&pl());){var it=H.callback;if(typeof it=="function"){H.callback=null,L=H.priorityLevel;var dt=it(H.expirationTime<=x);if(x=r.unstable_now(),typeof dt=="function"){H.callback=dt,_t(x),C=!0;break l}H===S(N)&&f(N),_t(x)}else f(N);H=S(N)}if(H!==null)C=!0;else{var o=S(p);o!==null&&ll(nt,o.startTime-x),C=!1}}break t}finally{H=null,L=Z,ot=!1}C=void 0}}finally{C?I():Ot=!1}}}var I;if(typeof zt=="function")I=function(){zt(Pt)};else if(typeof MessageChannel<"u"){var Rl=new MessageChannel,tl=Rl.port2;Rl.port1.onmessage=Pt,I=function(){tl.postMessage(null)}}else I=function(){V(Pt,0)};function ll(x,C){J=V(function(){x(r.unstable_now())},C)}r.unstable_IdlePriority=5,r.unstable_ImmediatePriority=1,r.unstable_LowPriority=4,r.unstable_NormalPriority=3,r.unstable_Profiling=null,r.unstable_UserBlockingPriority=2,r.unstable_cancelCallback=function(x){x.callback=null},r.unstable_forceFrameRate=function(x){0>x||125it?(x.sortIndex=Z,v(p,x),S(N)===null&&x===S(p)&&(G?(gt(J),J=-1):G=!0,ll(nt,Z-it))):(x.sortIndex=dt,v(N,x),ct||ot||(ct=!0,Ot||(Ot=!0,I()))),x},r.unstable_shouldYield=pl,r.unstable_wrapCallback=function(x){var C=L;return function(){var Z=L;L=C;try{return x.apply(this,arguments)}finally{L=Z}}}})(Af)),Af}var wd;function qm(){return wd||(wd=1,Tf.exports=Bm()),Tf.exports}var Ef={exports:{}},Wt={};var Ld;function Ym(){if(Ld)return Wt;Ld=1;var r=Rf();function v(N){var p="https://react.dev/errors/"+N;if(1"u"||typeof __REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE!="function"))try{__REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE(r)}catch(v){console.error(v)}}return r(),Ef.exports=Ym(),Ef.exports}var Kd;function Xm(){if(Kd)return Uu;Kd=1;var r=qm(),v=Rf(),S=Gm();function f(t){var l="https://react.dev/errors/"+t;if(1dt||(t.current=it[dt],it[dt]=null,dt--)}function j(t,l){dt++,it[dt]=t.current,t.current=l}var q=o(null),k=o(null),P=o(null),yt=o(null);function $t(t,l){switch(j(P,l),j(k,t),j(q,null),l.nodeType){case 9:case 11:t=(t=l.documentElement)&&(t=t.namespaceURI)?cd(t):0;break;default:if(t=l.tagName,l=l.namespaceURI)l=cd(l),t=fd(l,t);else switch(t){case"svg":t=1;break;case"math":t=2;break;default:t=0}}M(q),j(q,t)}function Ct(){M(q),M(k),M(P)}function qa(t){t.memoizedState!==null&&j(yt,t);var l=q.current,e=fd(l,t.type);l!==e&&(j(k,t),j(q,e))}function Hu(t){k.current===t&&(M(q),M(k)),yt.current===t&&(M(yt),Mu._currentValue=Z)}var li,Bf;function Ue(t){if(li===void 0)try{throw Error()}catch(e){var l=e.stack.trim().match(/\n( *(at )?)/);li=l&&l[1]||"",Bf=-1)":-1u||o[a]!==h[u]){var z=` -`+o[a].replace(" at new "," at ");return t.displayName&&z.includes("")&&(z=z.replace("",t.displayName)),z}while(1<=a&&0<=u);break}}}finally{ei=!1,Error.prepareStackTrace=e}return(e=t?t.displayName||t.name:"")?Ue(e):""}function o0(t,l){switch(t.tag){case 26:case 27:case 5:return Ue(t.type);case 16:return Ue("Lazy");case 13:return t.child!==l&&l!==null?Ue("Suspense Fallback"):Ue("Suspense");case 19:return Ue("SuspenseList");case 0:case 15:return ai(t.type,!1);case 11:return ai(t.type.render,!1);case 1:return ai(t.type,!0);case 31:return Ue("Activity");default:return""}}function qf(t){try{var l="",e=null;do l+=o0(t,e),e=t,t=t.return;while(t);return l}catch(a){return` +`);for(u=a=0;au||s[a]!==h[u]){var z=` +`+s[a].replace(" at new "," at ");return t.displayName&&z.includes("")&&(z=z.replace("",t.displayName)),z}while(1<=a&&0<=u);break}}}finally{ei=!1,Error.prepareStackTrace=e}return(e=t?t.displayName||t.name:"")?Ue(e):""}function s0(t,l){switch(t.tag){case 26:case 27:case 5:return Ue(t.type);case 16:return Ue("Lazy");case 13:return t.child!==l&&l!==null?Ue("Suspense Fallback"):Ue("Suspense");case 19:return Ue("SuspenseList");case 0:case 15:return ai(t.type,!1);case 11:return ai(t.type.render,!1);case 1:return ai(t.type,!0);case 31:return Ue("Activity");default:return""}}function qf(t){try{var l="",e=null;do l+=s0(t,e),e=t,t=t.return;while(t);return l}catch(a){return` Error generating stack: `+a.message+` -`+a.stack}}var ui=Object.prototype.hasOwnProperty,ni=r.unstable_scheduleCallback,ii=r.unstable_cancelCallback,s0=r.unstable_shouldYield,d0=r.unstable_requestPaint,rl=r.unstable_now,y0=r.unstable_getCurrentPriorityLevel,Yf=r.unstable_ImmediatePriority,Gf=r.unstable_UserBlockingPriority,Bu=r.unstable_NormalPriority,m0=r.unstable_LowPriority,Xf=r.unstable_IdlePriority,h0=r.log,g0=r.unstable_setDisableYieldValue,Ya=null,ol=null;function ue(t){if(typeof h0=="function"&&g0(t),ol&&typeof ol.setStrictMode=="function")try{ol.setStrictMode(Ya,t)}catch{}}var sl=Math.clz32?Math.clz32:p0,v0=Math.log,b0=Math.LN2;function p0(t){return t>>>=0,t===0?32:31-(v0(t)/b0|0)|0}var qu=256,Yu=262144,Gu=4194304;function Ce(t){var l=t&42;if(l!==0)return l;switch(t&-t){case 1:return 1;case 2:return 2;case 4:return 4;case 8:return 8;case 16:return 16;case 32:return 32;case 64:return 64;case 128:return 128;case 256:case 512:case 1024:case 2048:case 4096:case 8192:case 16384:case 32768:case 65536:case 131072:return t&261888;case 262144:case 524288:case 1048576:case 2097152:return t&3932160;case 4194304:case 8388608:case 16777216:case 33554432:return t&62914560;case 67108864:return 67108864;case 134217728:return 134217728;case 268435456:return 268435456;case 536870912:return 536870912;case 1073741824:return 0;default:return t}}function Xu(t,l,e){var a=t.pendingLanes;if(a===0)return 0;var u=0,n=t.suspendedLanes,i=t.pingedLanes;t=t.warmLanes;var c=a&134217727;return c!==0?(a=c&~n,a!==0?u=Ce(a):(i&=c,i!==0?u=Ce(i):e||(e=c&~t,e!==0&&(u=Ce(e))))):(c=a&~n,c!==0?u=Ce(c):i!==0?u=Ce(i):e||(e=a&~t,e!==0&&(u=Ce(e)))),u===0?0:l!==0&&l!==u&&(l&n)===0&&(n=u&-u,e=l&-l,n>=e||n===32&&(e&4194048)!==0)?l:u}function Ga(t,l){return(t.pendingLanes&~(t.suspendedLanes&~t.pingedLanes)&l)===0}function S0(t,l){switch(t){case 1:case 2:case 4:case 8:case 64:return l+250;case 16:case 32:case 128:case 256:case 512:case 1024:case 2048:case 4096:case 8192:case 16384:case 32768:case 65536:case 131072:case 262144:case 524288:case 1048576:case 2097152:return l+5e3;case 4194304:case 8388608:case 16777216:case 33554432:return-1;case 67108864:case 134217728:case 268435456:case 536870912:case 1073741824:return-1;default:return-1}}function Qf(){var t=Gu;return Gu<<=1,(Gu&62914560)===0&&(Gu=4194304),t}function ci(t){for(var l=[],e=0;31>e;e++)l.push(t);return l}function Xa(t,l){t.pendingLanes|=l,l!==268435456&&(t.suspendedLanes=0,t.pingedLanes=0,t.warmLanes=0)}function x0(t,l,e,a,u,n){var i=t.pendingLanes;t.pendingLanes=e,t.suspendedLanes=0,t.pingedLanes=0,t.warmLanes=0,t.expiredLanes&=e,t.entangledLanes&=e,t.errorRecoveryDisabledLanes&=e,t.shellSuspendCounter=0;var c=t.entanglements,o=t.expirationTimes,h=t.hiddenUpdates;for(e=i&~e;0"u")return null;try{return t.activeElement||t.body}catch{return t.body}}var _0=/[\n"\\]/g;function Sl(t){return t.replace(_0,function(l){return"\\"+l.charCodeAt(0).toString(16)+" "})}function yi(t,l,e,a,u,n,i,c){t.name="",i!=null&&typeof i!="function"&&typeof i!="symbol"&&typeof i!="boolean"?t.type=i:t.removeAttribute("type"),l!=null?i==="number"?(l===0&&t.value===""||t.value!=l)&&(t.value=""+pl(l)):t.value!==""+pl(l)&&(t.value=""+pl(l)):i!=="submit"&&i!=="reset"||t.removeAttribute("value"),l!=null?mi(t,i,pl(l)):e!=null?mi(t,i,pl(e)):a!=null&&t.removeAttribute("value"),u==null&&n!=null&&(t.defaultChecked=!!n),u!=null&&(t.checked=u&&typeof u!="function"&&typeof u!="symbol"),c!=null&&typeof c!="function"&&typeof c!="symbol"&&typeof c!="boolean"?t.name=""+pl(c):t.removeAttribute("name")}function tr(t,l,e,a,u,n,i,c){if(n!=null&&typeof n!="function"&&typeof n!="symbol"&&typeof n!="boolean"&&(t.type=n),l!=null||e!=null){if(!(n!=="submit"&&n!=="reset"||l!=null)){di(t);return}e=e!=null?""+pl(e):"",l=l!=null?""+pl(l):e,c||l===t.value||(t.value=l),t.defaultValue=l}a=a??u,a=typeof a!="function"&&typeof a!="symbol"&&!!a,t.checked=c?t.checked:!!a,t.defaultChecked=!!a,i!=null&&typeof i!="function"&&typeof i!="symbol"&&typeof i!="boolean"&&(t.name=i),di(t)}function mi(t,l,e){l==="number"&&wu(t.ownerDocument)===t||t.defaultValue===""+e||(t.defaultValue=""+e)}function ea(t,l,e,a){if(t=t.options,l){l={};for(var u=0;u"u"||typeof window.document>"u"||typeof window.document.createElement>"u"),pi=!1;if(Ql)try{var La={};Object.defineProperty(La,"passive",{get:function(){pi=!0}}),window.addEventListener("test",La,La),window.removeEventListener("test",La,La)}catch{pi=!1}var ie=null,Si=null,Vu=null;function cr(){if(Vu)return Vu;var t,l=Si,e=l.length,a,u="value"in ie?ie.value:ie.textContent,n=u.length;for(t=0;t=Ja),yr=" ",mr=!1;function hr(t,l){switch(t){case"keyup":return ly.indexOf(l.keyCode)!==-1;case"keydown":return l.keyCode!==229;case"keypress":case"mousedown":case"focusout":return!0;default:return!1}}function gr(t){return t=t.detail,typeof t=="object"&&"data"in t?t.data:null}var ia=!1;function ay(t,l){switch(t){case"compositionend":return gr(l);case"keypress":return l.which!==32?null:(mr=!0,yr);case"textInput":return t=l.data,t===yr&&mr?null:t;default:return null}}function uy(t,l){if(ia)return t==="compositionend"||!Ai&&hr(t,l)?(t=cr(),Vu=Si=ie=null,ia=!1,t):null;switch(t){case"paste":return null;case"keypress":if(!(l.ctrlKey||l.altKey||l.metaKey)||l.ctrlKey&&l.altKey){if(l.char&&1=l)return{node:e,offset:l-t};t=a}t:{for(;e;){if(e.nextSibling){e=e.nextSibling;break t}e=e.parentNode}e=void 0}e=Er(e)}}function Mr(t,l){return t&&l?t===l?!0:t&&t.nodeType===3?!1:l&&l.nodeType===3?Mr(t,l.parentNode):"contains"in t?t.contains(l):t.compareDocumentPosition?!!(t.compareDocumentPosition(l)&16):!1:!1}function _r(t){t=t!=null&&t.ownerDocument!=null&&t.ownerDocument.defaultView!=null?t.ownerDocument.defaultView:window;for(var l=wu(t.document);l instanceof t.HTMLIFrameElement;){try{var e=typeof l.contentWindow.location.href=="string"}catch{e=!1}if(e)t=l.contentWindow;else break;l=wu(t.document)}return l}function Oi(t){var l=t&&t.nodeName&&t.nodeName.toLowerCase();return l&&(l==="input"&&(t.type==="text"||t.type==="search"||t.type==="tel"||t.type==="url"||t.type==="password")||l==="textarea"||t.contentEditable==="true")}var dy=Ql&&"documentMode"in document&&11>=document.documentMode,ca=null,Ni=null,Fa=null,Di=!1;function Or(t,l,e){var a=e.window===e?e.document:e.nodeType===9?e:e.ownerDocument;Di||ca==null||ca!==wu(a)||(a=ca,"selectionStart"in a&&Oi(a)?a={start:a.selectionStart,end:a.selectionEnd}:(a=(a.ownerDocument&&a.ownerDocument.defaultView||window).getSelection(),a={anchorNode:a.anchorNode,anchorOffset:a.anchorOffset,focusNode:a.focusNode,focusOffset:a.focusOffset}),Fa&&$a(Fa,a)||(Fa=a,a=Gn(Ni,"onSelect"),0>=i,u-=i,Hl=1<<32-sl(l)+u|e<$?(at=Y,Y=null):at=Y.sibling;var rt=g(y,Y,m[$],T);if(rt===null){Y===null&&(Y=at);break}t&&Y&&rt.alternate===null&&l(y,Y),d=n(rt,d,$),ft===null?X=rt:ft.sibling=rt,ft=rt,Y=at}if($===m.length)return e(y,Y),ut&&wl(y,$),X;if(Y===null){for(;$$?(at=Y,Y=null):at=Y.sibling;var Oe=g(y,Y,rt.value,T);if(Oe===null){Y===null&&(Y=at);break}t&&Y&&Oe.alternate===null&&l(y,Y),d=n(Oe,d,$),ft===null?X=Oe:ft.sibling=Oe,ft=Oe,Y=at}if(rt.done)return e(y,Y),ut&&wl(y,$),X;if(Y===null){for(;!rt.done;$++,rt=m.next())rt=E(y,rt.value,T),rt!==null&&(d=n(rt,d,$),ft===null?X=rt:ft.sibling=rt,ft=rt);return ut&&wl(y,$),X}for(Y=a(Y);!rt.done;$++,rt=m.next())rt=b(Y,y,$,rt.value,T),rt!==null&&(t&&rt.alternate!==null&&Y.delete(rt.key===null?$:rt.key),d=n(rt,d,$),ft===null?X=rt:ft.sibling=rt,ft=rt);return t&&Y.forEach(function(Cm){return l(y,Cm)}),ut&&wl(y,$),X}function pt(y,d,m,T){if(typeof m=="object"&&m!==null&&m.type===G&&m.key===null&&(m=m.props.children),typeof m=="object"&&m!==null){switch(m.$$typeof){case st:t:{for(var X=m.key;d!==null;){if(d.key===X){if(X=m.type,X===G){if(d.tag===7){e(y,d.sibling),T=u(d,m.props.children),T.return=y,y=T;break t}}else if(d.elementType===X||typeof X=="object"&&X!==null&&X.$$typeof===Rt&&we(X)===d.type){e(y,d.sibling),T=u(d,m.props),au(T,m),T.return=y,y=T;break t}e(y,d);break}else l(y,d);d=d.sibling}m.type===G?(T=Ye(m.props.children,y.mode,T,m.key),T.return=y,y=T):(T=ln(m.type,m.key,m.props,null,y.mode,T),au(T,m),T.return=y,y=T)}return i(y);case ct:t:{for(X=m.key;d!==null;){if(d.key===X)if(d.tag===4&&d.stateNode.containerInfo===m.containerInfo&&d.stateNode.implementation===m.implementation){e(y,d.sibling),T=u(d,m.children||[]),T.return=y,y=T;break t}else{e(y,d);break}else l(y,d);d=d.sibling}T=qi(m,y.mode,T),T.return=y,y=T}return i(y);case Rt:return m=we(m),pt(y,d,m,T)}if(ll(m))return B(y,d,m,T);if(I(m)){if(X=I(m),typeof X!="function")throw Error(f(150));return m=X.call(m),w(y,d,m,T)}if(typeof m.then=="function")return pt(y,d,rn(m),T);if(m.$$typeof===zt)return pt(y,d,un(y,m),T);on(y,m)}return typeof m=="string"&&m!==""||typeof m=="number"||typeof m=="bigint"?(m=""+m,d!==null&&d.tag===6?(e(y,d.sibling),T=u(d,m),T.return=y,y=T):(e(y,d),T=Bi(m,y.mode,T),T.return=y,y=T),i(y)):e(y,d)}return function(y,d,m,T){try{eu=0;var X=pt(y,d,m,T);return ba=null,X}catch(Y){if(Y===va||Y===cn)throw Y;var ft=yl(29,Y,null,y.mode);return ft.lanes=T,ft.return=y,ft}}}var Ve=Fr(!0),Ir=Fr(!1),se=!1;function Wi(t){t.updateQueue={baseState:t.memoizedState,firstBaseUpdate:null,lastBaseUpdate:null,shared:{pending:null,lanes:0,hiddenCallbacks:null},callbacks:null}}function $i(t,l){t=t.updateQueue,l.updateQueue===t&&(l.updateQueue={baseState:t.baseState,firstBaseUpdate:t.firstBaseUpdate,lastBaseUpdate:t.lastBaseUpdate,shared:t.shared,callbacks:null})}function de(t){return{lane:t,tag:0,payload:null,callback:null,next:null}}function ye(t,l,e){var a=t.updateQueue;if(a===null)return null;if(a=a.shared,(ot&2)!==0){var u=a.pending;return u===null?l.next=l:(l.next=u.next,u.next=l),a.pending=l,l=tn(t),Hr(t,null,e),l}return Pu(t,a,l,e),tn(t)}function uu(t,l,e){if(l=l.updateQueue,l!==null&&(l=l.shared,(e&4194048)!==0)){var a=l.lanes;a&=t.pendingLanes,e|=a,l.lanes=e,wf(t,e)}}function Fi(t,l){var e=t.updateQueue,a=t.alternate;if(a!==null&&(a=a.updateQueue,e===a)){var u=null,n=null;if(e=e.firstBaseUpdate,e!==null){do{var i={lane:e.lane,tag:e.tag,payload:e.payload,callback:null,next:null};n===null?u=n=i:n=n.next=i,e=e.next}while(e!==null);n===null?u=n=l:n=n.next=l}else u=n=l;e={baseState:a.baseState,firstBaseUpdate:u,lastBaseUpdate:n,shared:a.shared,callbacks:a.callbacks},t.updateQueue=e;return}t=e.lastBaseUpdate,t===null?e.firstBaseUpdate=l:t.next=l,e.lastBaseUpdate=l}var Ii=!1;function nu(){if(Ii){var t=ga;if(t!==null)throw t}}function iu(t,l,e,a){Ii=!1;var u=t.updateQueue;se=!1;var n=u.firstBaseUpdate,i=u.lastBaseUpdate,c=u.shared.pending;if(c!==null){u.shared.pending=null;var o=c,h=o.next;o.next=null,i===null?n=h:i.next=h,i=o;var z=t.alternate;z!==null&&(z=z.updateQueue,c=z.lastBaseUpdate,c!==i&&(c===null?z.firstBaseUpdate=h:c.next=h,z.lastBaseUpdate=o))}if(n!==null){var E=u.baseState;i=0,z=h=o=null,c=n;do{var g=c.lane&-536870913,b=g!==c.lane;if(b?(et&g)===g:(a&g)===g){g!==0&&g===ha&&(Ii=!0),z!==null&&(z=z.next={lane:0,tag:c.tag,payload:c.payload,callback:null,next:null});t:{var B=t,w=c;g=l;var pt=e;switch(w.tag){case 1:if(B=w.payload,typeof B=="function"){E=B.call(pt,E,g);break t}E=B;break t;case 3:B.flags=B.flags&-65537|128;case 0:if(B=w.payload,g=typeof B=="function"?B.call(pt,E,g):B,g==null)break t;E=H({},E,g);break t;case 2:se=!0}}g=c.callback,g!==null&&(t.flags|=64,b&&(t.flags|=8192),b=u.callbacks,b===null?u.callbacks=[g]:b.push(g))}else b={lane:g,tag:c.tag,payload:c.payload,callback:c.callback,next:null},z===null?(h=z=b,o=E):z=z.next=b,i|=g;if(c=c.next,c===null){if(c=u.shared.pending,c===null)break;b=c,c=b.next,b.next=null,u.lastBaseUpdate=b,u.shared.pending=null}}while(!0);z===null&&(o=E),u.baseState=o,u.firstBaseUpdate=h,u.lastBaseUpdate=z,n===null&&(u.shared.lanes=0),be|=i,t.lanes=i,t.memoizedState=E}}function Pr(t,l){if(typeof t!="function")throw Error(f(191,t));t.call(l)}function to(t,l){var e=t.callbacks;if(e!==null)for(t.callbacks=null,t=0;tn?n:8;var i=x.T,c={};x.T=c,vc(t,!1,l,e);try{var o=u(),h=x.S;if(h!==null&&h(c,o),o!==null&&typeof o=="object"&&typeof o.then=="function"){var z=xy(o,a);ru(t,l,z,bl(t))}else ru(t,l,a,bl(t))}catch(E){ru(t,l,{then:function(){},status:"rejected",reason:E},bl())}finally{C.p=n,i!==null&&c.types!==null&&(i.types=c.types),x.T=i}}function _y(){}function hc(t,l,e,a){if(t.tag!==5)throw Error(f(476));var u=jo(t).queue;Co(t,u,l,Z,e===null?_y:function(){return Ro(t),e(a)})}function jo(t){var l=t.memoizedState;if(l!==null)return l;l={memoizedState:Z,baseState:Z,baseQueue:null,queue:{pending:null,lanes:0,dispatch:null,lastRenderedReducer:Jl,lastRenderedState:Z},next:null};var e={};return l.next={memoizedState:e,baseState:e,baseQueue:null,queue:{pending:null,lanes:0,dispatch:null,lastRenderedReducer:Jl,lastRenderedState:e},next:null},t.memoizedState=l,t=t.alternate,t!==null&&(t.memoizedState=l),l}function Ro(t){var l=jo(t);l.next===null&&(l=t.alternate.memoizedState),ru(t,l.next.queue,{},bl())}function gc(){return Vt(Mu)}function Ho(){return jt().memoizedState}function Bo(){return jt().memoizedState}function Oy(t){for(var l=t.return;l!==null;){switch(l.tag){case 24:case 3:var e=bl();t=de(e);var a=ye(l,t,e);a!==null&&(fl(a,l,e),uu(a,l,e)),l={cache:Vi()},t.payload=l;return}l=l.return}}function Ny(t,l,e){var a=bl();e={lane:a,revertLane:0,gesture:null,action:e,hasEagerState:!1,eagerState:null,next:null},Sn(t)?Yo(l,e):(e=Ri(t,l,e,a),e!==null&&(fl(e,t,a),Go(e,l,a)))}function qo(t,l,e){var a=bl();ru(t,l,e,a)}function ru(t,l,e,a){var u={lane:a,revertLane:0,gesture:null,action:e,hasEagerState:!1,eagerState:null,next:null};if(Sn(t))Yo(l,u);else{var n=t.alternate;if(t.lanes===0&&(n===null||n.lanes===0)&&(n=l.lastRenderedReducer,n!==null))try{var i=l.lastRenderedState,c=n(i,e);if(u.hasEagerState=!0,u.eagerState=c,dl(c,i))return Pu(t,l,u,0),St===null&&Iu(),!1}catch{}if(e=Ri(t,l,u,a),e!==null)return fl(e,t,a),Go(e,l,a),!0}return!1}function vc(t,l,e,a){if(a={lane:2,revertLane:Wc(),gesture:null,action:a,hasEagerState:!1,eagerState:null,next:null},Sn(t)){if(l)throw Error(f(479))}else l=Ri(t,e,a,2),l!==null&&fl(l,t,2)}function Sn(t){var l=t.alternate;return t===W||l!==null&&l===W}function Yo(t,l){Sa=yn=!0;var e=t.pending;e===null?l.next=l:(l.next=e.next,e.next=l),t.pending=l}function Go(t,l,e){if((e&4194048)!==0){var a=l.lanes;a&=t.pendingLanes,e|=a,l.lanes=e,wf(t,e)}}var ou={readContext:Vt,use:gn,useCallback:Nt,useContext:Nt,useEffect:Nt,useImperativeHandle:Nt,useLayoutEffect:Nt,useInsertionEffect:Nt,useMemo:Nt,useReducer:Nt,useRef:Nt,useState:Nt,useDebugValue:Nt,useDeferredValue:Nt,useTransition:Nt,useSyncExternalStore:Nt,useId:Nt,useHostTransitionStatus:Nt,useFormState:Nt,useActionState:Nt,useOptimistic:Nt,useMemoCache:Nt,useCacheRefresh:Nt};ou.useEffectEvent=Nt;var Xo={readContext:Vt,use:gn,useCallback:function(t,l){return $t().memoizedState=[t,l===void 0?null:l],t},useContext:Vt,useEffect:To,useImperativeHandle:function(t,l,e){e=e!=null?e.concat([t]):null,bn(4194308,4,_o.bind(null,l,t),e)},useLayoutEffect:function(t,l){return bn(4194308,4,t,l)},useInsertionEffect:function(t,l){bn(4,2,t,l)},useMemo:function(t,l){var e=$t();l=l===void 0?null:l;var a=t();if(Ke){ue(!0);try{t()}finally{ue(!1)}}return e.memoizedState=[a,l],a},useReducer:function(t,l,e){var a=$t();if(e!==void 0){var u=e(l);if(Ke){ue(!0);try{e(l)}finally{ue(!1)}}}else u=l;return a.memoizedState=a.baseState=u,t={pending:null,lanes:0,dispatch:null,lastRenderedReducer:t,lastRenderedState:u},a.queue=t,t=t.dispatch=Ny.bind(null,W,t),[a.memoizedState,t]},useRef:function(t){var l=$t();return t={current:t},l.memoizedState=t},useState:function(t){t=oc(t);var l=t.queue,e=qo.bind(null,W,l);return l.dispatch=e,[t.memoizedState,e]},useDebugValue:yc,useDeferredValue:function(t,l){var e=$t();return mc(e,t,l)},useTransition:function(){var t=oc(!1);return t=Co.bind(null,W,t.queue,!0,!1),$t().memoizedState=t,[!1,t]},useSyncExternalStore:function(t,l,e){var a=W,u=$t();if(ut){if(e===void 0)throw Error(f(407));e=e()}else{if(e=l(),St===null)throw Error(f(349));(et&127)!==0||io(a,l,e)}u.memoizedState=e;var n={value:e,getSnapshot:l};return u.queue=n,To(fo.bind(null,a,n,t),[t]),a.flags|=2048,za(9,{destroy:void 0},co.bind(null,a,n,e,l),null),e},useId:function(){var t=$t(),l=St.identifierPrefix;if(ut){var e=Bl,a=Hl;e=(a&~(1<<32-sl(a)-1)).toString(32)+e,l="_"+l+"R_"+e,e=mn++,0<\/script>",n=n.removeChild(n.firstChild);break;case"select":n=typeof a.is=="string"?i.createElement("select",{is:a.is}):i.createElement("select"),a.multiple?n.multiple=!0:a.size&&(n.size=a.size);break;default:n=typeof a.is=="string"?i.createElement(u,{is:a.is}):i.createElement(u)}}n[wt]=l,n[el]=a;t:for(i=l.child;i!==null;){if(i.tag===5||i.tag===6)n.appendChild(i.stateNode);else if(i.tag!==4&&i.tag!==27&&i.child!==null){i.child.return=i,i=i.child;continue}if(i===l)break t;for(;i.sibling===null;){if(i.return===null||i.return===l)break t;i=i.return}i.sibling.return=i.return,i=i.sibling}l.stateNode=n;t:switch(Jt(n,u,a),u){case"button":case"input":case"select":case"textarea":a=!!a.autoFocus;break t;case"img":a=!0;break t;default:a=!1}a&&Wl(l)}}return Et(l),Uc(l,l.type,t===null?null:t.memoizedProps,l.pendingProps,e),null;case 6:if(t&&l.stateNode!=null)t.memoizedProps!==a&&Wl(l);else{if(typeof a!="string"&&l.stateNode===null)throw Error(f(166));if(t=P.current,ya(l)){if(t=l.stateNode,e=l.memoizedProps,a=null,u=Lt,u!==null)switch(u.tag){case 27:case 5:a=u.memoizedProps}t[wt]=l,t=!!(t.nodeValue===e||a!==null&&a.suppressHydrationWarning===!0||nd(t.nodeValue,e)),t||re(l,!0)}else t=Xn(t).createTextNode(a),t[wt]=l,l.stateNode=t}return Et(l),null;case 31:if(e=l.memoizedState,t===null||t.memoizedState!==null){if(a=ya(l),e!==null){if(t===null){if(!a)throw Error(f(318));if(t=l.memoizedState,t=t!==null?t.dehydrated:null,!t)throw Error(f(557));t[wt]=l}else Ge(),(l.flags&128)===0&&(l.memoizedState=null),l.flags|=4;Et(l),t=!1}else e=Qi(),t!==null&&t.memoizedState!==null&&(t.memoizedState.hydrationErrors=e),t=!0;if(!t)return l.flags&256?(hl(l),l):(hl(l),null);if((l.flags&128)!==0)throw Error(f(558))}return Et(l),null;case 13:if(a=l.memoizedState,t===null||t.memoizedState!==null&&t.memoizedState.dehydrated!==null){if(u=ya(l),a!==null&&a.dehydrated!==null){if(t===null){if(!u)throw Error(f(318));if(u=l.memoizedState,u=u!==null?u.dehydrated:null,!u)throw Error(f(317));u[wt]=l}else Ge(),(l.flags&128)===0&&(l.memoizedState=null),l.flags|=4;Et(l),u=!1}else u=Qi(),t!==null&&t.memoizedState!==null&&(t.memoizedState.hydrationErrors=u),u=!0;if(!u)return l.flags&256?(hl(l),l):(hl(l),null)}return hl(l),(l.flags&128)!==0?(l.lanes=e,l):(e=a!==null,t=t!==null&&t.memoizedState!==null,e&&(a=l.child,u=null,a.alternate!==null&&a.alternate.memoizedState!==null&&a.alternate.memoizedState.cachePool!==null&&(u=a.alternate.memoizedState.cachePool.pool),n=null,a.memoizedState!==null&&a.memoizedState.cachePool!==null&&(n=a.memoizedState.cachePool.pool),n!==u&&(a.flags|=2048)),e!==t&&e&&(l.child.flags|=8192),An(l,l.updateQueue),Et(l),null);case 4:return Ut(),t===null&&Pc(l.stateNode.containerInfo),Et(l),null;case 10:return Vl(l.type),Et(l),null;case 19:if(M(Ct),a=l.memoizedState,a===null)return Et(l),null;if(u=(l.flags&128)!==0,n=a.rendering,n===null)if(u)du(a,!1);else{if(Dt!==0||t!==null&&(t.flags&128)!==0)for(t=l.child;t!==null;){if(n=dn(t),n!==null){for(l.flags|=128,du(a,!1),t=n.updateQueue,l.updateQueue=t,An(l,t),l.subtreeFlags=0,t=e,e=l.child;e!==null;)Br(e,t),e=e.sibling;return j(Ct,Ct.current&1|2),ut&&wl(l,a.treeForkCount),l.child}t=t.sibling}a.tail!==null&&rl()>Dn&&(l.flags|=128,u=!0,du(a,!1),l.lanes=4194304)}else{if(!u)if(t=dn(n),t!==null){if(l.flags|=128,u=!0,t=t.updateQueue,l.updateQueue=t,An(l,t),du(a,!0),a.tail===null&&a.tailMode==="hidden"&&!n.alternate&&!ut)return Et(l),null}else 2*rl()-a.renderingStartTime>Dn&&e!==536870912&&(l.flags|=128,u=!0,du(a,!1),l.lanes=4194304);a.isBackwards?(n.sibling=l.child,l.child=n):(t=a.last,t!==null?t.sibling=n:l.child=n,a.last=n)}return a.tail!==null?(t=a.tail,a.rendering=t,a.tail=t.sibling,a.renderingStartTime=rl(),t.sibling=null,e=Ct.current,j(Ct,u?e&1|2:e&1),ut&&wl(l,a.treeForkCount),t):(Et(l),null);case 22:case 23:return hl(l),tc(),a=l.memoizedState!==null,t!==null?t.memoizedState!==null!==a&&(l.flags|=8192):a&&(l.flags|=8192),a?(e&536870912)!==0&&(l.flags&128)===0&&(Et(l),l.subtreeFlags&6&&(l.flags|=8192)):Et(l),e=l.updateQueue,e!==null&&An(l,e.retryQueue),e=null,t!==null&&t.memoizedState!==null&&t.memoizedState.cachePool!==null&&(e=t.memoizedState.cachePool.pool),a=null,l.memoizedState!==null&&l.memoizedState.cachePool!==null&&(a=l.memoizedState.cachePool.pool),a!==e&&(l.flags|=2048),t!==null&&M(Ze),null;case 24:return e=null,t!==null&&(e=t.memoizedState.cache),l.memoizedState.cache!==e&&(l.flags|=2048),Vl(Ht),Et(l),null;case 25:return null;case 30:return null}throw Error(f(156,l.tag))}function Ry(t,l){switch(Gi(l),l.tag){case 1:return t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 3:return Vl(Ht),Ut(),t=l.flags,(t&65536)!==0&&(t&128)===0?(l.flags=t&-65537|128,l):null;case 26:case 27:case 5:return Hu(l),null;case 31:if(l.memoizedState!==null){if(hl(l),l.alternate===null)throw Error(f(340));Ge()}return t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 13:if(hl(l),t=l.memoizedState,t!==null&&t.dehydrated!==null){if(l.alternate===null)throw Error(f(340));Ge()}return t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 19:return M(Ct),null;case 4:return Ut(),null;case 10:return Vl(l.type),null;case 22:case 23:return hl(l),tc(),t!==null&&M(Ze),t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 24:return Vl(Ht),null;case 25:return null;default:return null}}function os(t,l){switch(Gi(l),l.tag){case 3:Vl(Ht),Ut();break;case 26:case 27:case 5:Hu(l);break;case 4:Ut();break;case 31:l.memoizedState!==null&&hl(l);break;case 13:hl(l);break;case 19:M(Ct);break;case 10:Vl(l.type);break;case 22:case 23:hl(l),tc(),t!==null&&M(Ze);break;case 24:Vl(Ht)}}function yu(t,l){try{var e=l.updateQueue,a=e!==null?e.lastEffect:null;if(a!==null){var u=a.next;e=u;do{if((e.tag&t)===t){a=void 0;var n=e.create,i=e.inst;a=n(),i.destroy=a}e=e.next}while(e!==u)}}catch(c){ht(l,l.return,c)}}function ge(t,l,e){try{var a=l.updateQueue,u=a!==null?a.lastEffect:null;if(u!==null){var n=u.next;a=n;do{if((a.tag&t)===t){var i=a.inst,c=i.destroy;if(c!==void 0){i.destroy=void 0,u=l;var o=e,h=c;try{h()}catch(z){ht(u,o,z)}}}a=a.next}while(a!==n)}}catch(z){ht(l,l.return,z)}}function ss(t){var l=t.updateQueue;if(l!==null){var e=t.stateNode;try{to(l,e)}catch(a){ht(t,t.return,a)}}}function ds(t,l,e){e.props=Je(t.type,t.memoizedProps),e.state=t.memoizedState;try{e.componentWillUnmount()}catch(a){ht(t,l,a)}}function mu(t,l){try{var e=t.ref;if(e!==null){switch(t.tag){case 26:case 27:case 5:var a=t.stateNode;break;case 30:a=t.stateNode;break;default:a=t.stateNode}typeof e=="function"?t.refCleanup=e(a):e.current=a}}catch(u){ht(t,l,u)}}function ql(t,l){var e=t.ref,a=t.refCleanup;if(e!==null)if(typeof a=="function")try{a()}catch(u){ht(t,l,u)}finally{t.refCleanup=null,t=t.alternate,t!=null&&(t.refCleanup=null)}else if(typeof e=="function")try{e(null)}catch(u){ht(t,l,u)}else e.current=null}function ys(t){var l=t.type,e=t.memoizedProps,a=t.stateNode;try{t:switch(l){case"button":case"input":case"select":case"textarea":e.autoFocus&&a.focus();break t;case"img":e.src?a.src=e.src:e.srcSet&&(a.srcset=e.srcSet)}}catch(u){ht(t,t.return,u)}}function Cc(t,l,e){try{var a=t.stateNode;em(a,t.type,e,l),a[el]=l}catch(u){ht(t,t.return,u)}}function ms(t){return t.tag===5||t.tag===3||t.tag===26||t.tag===27&&Te(t.type)||t.tag===4}function jc(t){t:for(;;){for(;t.sibling===null;){if(t.return===null||ms(t.return))return null;t=t.return}for(t.sibling.return=t.return,t=t.sibling;t.tag!==5&&t.tag!==6&&t.tag!==18;){if(t.tag===27&&Te(t.type)||t.flags&2||t.child===null||t.tag===4)continue t;t.child.return=t,t=t.child}if(!(t.flags&2))return t.stateNode}}function Rc(t,l,e){var a=t.tag;if(a===5||a===6)t=t.stateNode,l?(e.nodeType===9?e.body:e.nodeName==="HTML"?e.ownerDocument.body:e).insertBefore(t,l):(l=e.nodeType===9?e.body:e.nodeName==="HTML"?e.ownerDocument.body:e,l.appendChild(t),e=e._reactRootContainer,e!=null||l.onclick!==null||(l.onclick=Xl));else if(a!==4&&(a===27&&Te(t.type)&&(e=t.stateNode,l=null),t=t.child,t!==null))for(Rc(t,l,e),t=t.sibling;t!==null;)Rc(t,l,e),t=t.sibling}function Mn(t,l,e){var a=t.tag;if(a===5||a===6)t=t.stateNode,l?e.insertBefore(t,l):e.appendChild(t);else if(a!==4&&(a===27&&Te(t.type)&&(e=t.stateNode),t=t.child,t!==null))for(Mn(t,l,e),t=t.sibling;t!==null;)Mn(t,l,e),t=t.sibling}function hs(t){var l=t.stateNode,e=t.memoizedProps;try{for(var a=t.type,u=l.attributes;u.length;)l.removeAttributeNode(u[0]);Jt(l,a,e),l[wt]=t,l[el]=e}catch(n){ht(t,t.return,n)}}var $l=!1,Yt=!1,Hc=!1,gs=typeof WeakSet=="function"?WeakSet:Set,Qt=null;function Hy(t,l){if(t=t.containerInfo,ef=Jn,t=_r(t),Oi(t)){if("selectionStart"in t)var e={start:t.selectionStart,end:t.selectionEnd};else t:{e=(e=t.ownerDocument)&&e.defaultView||window;var a=e.getSelection&&e.getSelection();if(a&&a.rangeCount!==0){e=a.anchorNode;var u=a.anchorOffset,n=a.focusNode;a=a.focusOffset;try{e.nodeType,n.nodeType}catch{e=null;break t}var i=0,c=-1,o=-1,h=0,z=0,E=t,g=null;l:for(;;){for(var b;E!==e||u!==0&&E.nodeType!==3||(c=i+u),E!==n||a!==0&&E.nodeType!==3||(o=i+a),E.nodeType===3&&(i+=E.nodeValue.length),(b=E.firstChild)!==null;)g=E,E=b;for(;;){if(E===t)break l;if(g===e&&++h===u&&(c=i),g===n&&++z===a&&(o=i),(b=E.nextSibling)!==null)break;E=g,g=E.parentNode}E=b}e=c===-1||o===-1?null:{start:c,end:o}}else e=null}e=e||{start:0,end:0}}else e=null;for(af={focusedElem:t,selectionRange:e},Jn=!1,Qt=l;Qt!==null;)if(l=Qt,t=l.child,(l.subtreeFlags&1028)!==0&&t!==null)t.return=l,Qt=t;else for(;Qt!==null;){switch(l=Qt,n=l.alternate,t=l.flags,l.tag){case 0:if((t&4)!==0&&(t=l.updateQueue,t=t!==null?t.events:null,t!==null))for(e=0;e title"))),Jt(n,a,e),n[wt]=t,Xt(n),a=n;break t;case"link":var i=zd("link","href",u).get(a+(e.href||""));if(i){for(var c=0;cpt&&(i=pt,pt=w,w=i);var y=Ar(c,w),d=Ar(c,pt);if(y&&d&&(b.rangeCount!==1||b.anchorNode!==y.node||b.anchorOffset!==y.offset||b.focusNode!==d.node||b.focusOffset!==d.offset)){var m=E.createRange();m.setStart(y.node,y.offset),b.removeAllRanges(),w>pt?(b.addRange(m),b.extend(d.node,d.offset)):(m.setEnd(d.node,d.offset),b.addRange(m))}}}}for(E=[],b=c;b=b.parentNode;)b.nodeType===1&&E.push({element:b,left:b.scrollLeft,top:b.scrollTop});for(typeof c.focus=="function"&&c.focus(),c=0;ce?32:e,x.T=null,e=Zc,Zc=null;var n=Se,i=le;if(Gt=0,_a=Se=null,le=0,(ot&6)!==0)throw Error(f(331));var c=ot;if(ot|=4,_s(n.current),Es(n,n.current,i,e),ot=c,Su(0,!1),ol&&typeof ol.onPostCommitFiberRoot=="function")try{ol.onPostCommitFiberRoot(Ya,n)}catch{}return!0}finally{C.p=u,x.T=a,Vs(t,l)}}function Js(t,l,e){l=zl(e,l),l=xc(t.stateNode,l,2),t=ye(t,l,2),t!==null&&(Xa(t,2),Yl(t))}function ht(t,l,e){if(t.tag===3)Js(t,t,e);else for(;l!==null;){if(l.tag===3){Js(l,t,e);break}else if(l.tag===1){var a=l.stateNode;if(typeof l.type.getDerivedStateFromError=="function"||typeof a.componentDidCatch=="function"&&(pe===null||!pe.has(a))){t=zl(e,t),e=ko(2),a=ye(l,e,2),a!==null&&(Wo(e,a,l,t),Xa(a,2),Yl(a));break}}l=l.return}}function Kc(t,l,e){var a=t.pingCache;if(a===null){a=t.pingCache=new Yy;var u=new Set;a.set(l,u)}else u=a.get(l),u===void 0&&(u=new Set,a.set(l,u));u.has(e)||(Yc=!0,u.add(e),t=wy.bind(null,t,l,e),l.then(t,t))}function wy(t,l,e){var a=t.pingCache;a!==null&&a.delete(l),t.pingedLanes|=t.suspendedLanes&e,t.warmLanes&=~e,St===t&&(et&e)===e&&(Dt===4||Dt===3&&(et&62914560)===et&&300>rl()-Nn?(ot&2)===0&&Oa(t,0):Gc|=e,Ma===et&&(Ma=0)),Yl(t)}function ks(t,l){l===0&&(l=Qf()),t=qe(t,l),t!==null&&(Xa(t,l),Yl(t))}function Ly(t){var l=t.memoizedState,e=0;l!==null&&(e=l.retryLane),ks(t,e)}function Vy(t,l){var e=0;switch(t.tag){case 31:case 13:var a=t.stateNode,u=t.memoizedState;u!==null&&(e=u.retryLane);break;case 19:a=t.stateNode;break;case 22:a=t.stateNode._retryCache;break;default:throw Error(f(314))}a!==null&&a.delete(l),ks(t,e)}function Ky(t,l){return ni(t,l)}var Bn=null,Da=null,Jc=!1,qn=!1,kc=!1,ze=0;function Yl(t){t!==Da&&t.next===null&&(Da===null?Bn=Da=t:Da=Da.next=t),qn=!0,Jc||(Jc=!0,ky())}function Su(t,l){if(!kc&&qn){kc=!0;do for(var e=!1,a=Bn;a!==null;){if(t!==0){var u=a.pendingLanes;if(u===0)var n=0;else{var i=a.suspendedLanes,c=a.pingedLanes;n=(1<<31-sl(42|t)+1)-1,n&=u&~(i&~c),n=n&201326741?n&201326741|1:n?n|2:0}n!==0&&(e=!0,Is(a,n))}else n=et,n=Xu(a,a===St?n:0,a.cancelPendingCommit!==null||a.timeoutHandle!==-1),(n&3)===0||Ga(a,n)||(e=!0,Is(a,n));a=a.next}while(e);kc=!1}}function Jy(){Ws()}function Ws(){qn=Jc=!1;var t=0;ze!==0&&um()&&(t=ze);for(var l=rl(),e=null,a=Bn;a!==null;){var u=a.next,n=$s(a,l);n===0?(a.next=null,e===null?Bn=u:e.next=u,u===null&&(Da=e)):(e=a,(t!==0||(n&3)!==0)&&(qn=!0)),a=u}Gt!==0&&Gt!==5||Su(t),ze!==0&&(ze=0)}function $s(t,l){for(var e=t.suspendedLanes,a=t.pingedLanes,u=t.expirationTimes,n=t.pendingLanes&-62914561;0c)break;var z=o.transferSize,E=o.initiatorType;z&&id(E)&&(o=o.responseEnd,i+=z*(o"u"?null:document;function bd(t,l,e){var a=Ua;if(a&&typeof l=="string"&&l){var u=Sl(l);u='link[rel="'+t+'"][href="'+u+'"]',typeof e=="string"&&(u+='[crossorigin="'+e+'"]'),vd.has(u)||(vd.add(u),t={rel:t,crossOrigin:e,href:l},a.querySelector(u)===null&&(l=a.createElement("link"),Jt(l,"link",t),Xt(l),a.head.appendChild(l)))}}function ym(t){ee.D(t),bd("dns-prefetch",t,null)}function mm(t,l){ee.C(t,l),bd("preconnect",t,l)}function hm(t,l,e){ee.L(t,l,e);var a=Ua;if(a&&t&&l){var u='link[rel="preload"][as="'+Sl(l)+'"]';l==="image"&&e&&e.imageSrcSet?(u+='[imagesrcset="'+Sl(e.imageSrcSet)+'"]',typeof e.imageSizes=="string"&&(u+='[imagesizes="'+Sl(e.imageSizes)+'"]')):u+='[href="'+Sl(t)+'"]';var n=u;switch(l){case"style":n=Ca(t);break;case"script":n=ja(t)}Ol.has(n)||(t=H({rel:"preload",href:l==="image"&&e&&e.imageSrcSet?void 0:t,as:l},e),Ol.set(n,t),a.querySelector(u)!==null||l==="style"&&a.querySelector(Eu(n))||l==="script"&&a.querySelector(Au(n))||(l=a.createElement("link"),Jt(l,"link",t),Xt(l),a.head.appendChild(l)))}}function gm(t,l){ee.m(t,l);var e=Ua;if(e&&t){var a=l&&typeof l.as=="string"?l.as:"script",u='link[rel="modulepreload"][as="'+Sl(a)+'"][href="'+Sl(t)+'"]',n=u;switch(a){case"audioworklet":case"paintworklet":case"serviceworker":case"sharedworker":case"worker":case"script":n=ja(t)}if(!Ol.has(n)&&(t=H({rel:"modulepreload",href:t},l),Ol.set(n,t),e.querySelector(u)===null)){switch(a){case"audioworklet":case"paintworklet":case"serviceworker":case"sharedworker":case"worker":case"script":if(e.querySelector(Au(n)))return}a=e.createElement("link"),Jt(a,"link",t),Xt(a),e.head.appendChild(a)}}}function vm(t,l,e){ee.S(t,l,e);var a=Ua;if(a&&t){var u=ta(a).hoistableStyles,n=Ca(t);l=l||"default";var i=u.get(n);if(!i){var c={loading:0,preload:null};if(i=a.querySelector(Eu(n)))c.loading=5;else{t=H({rel:"stylesheet",href:t,"data-precedence":l},e),(e=Ol.get(n))&&sf(t,e);var o=i=a.createElement("link");Xt(o),Jt(o,"link",t),o._p=new Promise(function(h,z){o.onload=h,o.onerror=z}),o.addEventListener("load",function(){c.loading|=1}),o.addEventListener("error",function(){c.loading|=2}),c.loading|=4,Zn(i,l,a)}i={type:"stylesheet",instance:i,count:1,state:c},u.set(n,i)}}}function bm(t,l){ee.X(t,l);var e=Ua;if(e&&t){var a=ta(e).hoistableScripts,u=ja(t),n=a.get(u);n||(n=e.querySelector(Au(u)),n||(t=H({src:t,async:!0},l),(l=Ol.get(u))&&df(t,l),n=e.createElement("script"),Xt(n),Jt(n,"link",t),e.head.appendChild(n)),n={type:"script",instance:n,count:1,state:null},a.set(u,n))}}function pm(t,l){ee.M(t,l);var e=Ua;if(e&&t){var a=ta(e).hoistableScripts,u=ja(t),n=a.get(u);n||(n=e.querySelector(Au(u)),n||(t=H({src:t,async:!0,type:"module"},l),(l=Ol.get(u))&&df(t,l),n=e.createElement("script"),Xt(n),Jt(n,"link",t),e.head.appendChild(n)),n={type:"script",instance:n,count:1,state:null},a.set(u,n))}}function pd(t,l,e,a){var u=(u=P.current)?Qn(u):null;if(!u)throw Error(f(446));switch(t){case"meta":case"title":return null;case"style":return typeof e.precedence=="string"&&typeof e.href=="string"?(l=Ca(e.href),e=ta(u).hoistableStyles,a=e.get(l),a||(a={type:"style",instance:null,count:0,state:null},e.set(l,a)),a):{type:"void",instance:null,count:0,state:null};case"link":if(e.rel==="stylesheet"&&typeof e.href=="string"&&typeof e.precedence=="string"){t=Ca(e.href);var n=ta(u).hoistableStyles,i=n.get(t);if(i||(u=u.ownerDocument||u,i={type:"stylesheet",instance:null,count:0,state:{loading:0,preload:null}},n.set(t,i),(n=u.querySelector(Eu(t)))&&!n._p&&(i.instance=n,i.state.loading=5),Ol.has(t)||(e={rel:"preload",as:"style",href:e.href,crossOrigin:e.crossOrigin,integrity:e.integrity,media:e.media,hrefLang:e.hrefLang,referrerPolicy:e.referrerPolicy},Ol.set(t,e),n||Sm(u,t,e,i.state))),l&&a===null)throw Error(f(528,""));return i}if(l&&a!==null)throw Error(f(529,""));return null;case"script":return l=e.async,e=e.src,typeof e=="string"&&l&&typeof l!="function"&&typeof l!="symbol"?(l=ja(e),e=ta(u).hoistableScripts,a=e.get(l),a||(a={type:"script",instance:null,count:0,state:null},e.set(l,a)),a):{type:"void",instance:null,count:0,state:null};default:throw Error(f(444,t))}}function Ca(t){return'href="'+Sl(t)+'"'}function Eu(t){return'link[rel="stylesheet"]['+t+"]"}function Sd(t){return H({},t,{"data-precedence":t.precedence,precedence:null})}function Sm(t,l,e,a){t.querySelector('link[rel="preload"][as="style"]['+l+"]")?a.loading=1:(l=t.createElement("link"),a.preload=l,l.addEventListener("load",function(){return a.loading|=1}),l.addEventListener("error",function(){return a.loading|=2}),Jt(l,"link",e),Xt(l),t.head.appendChild(l))}function ja(t){return'[src="'+Sl(t)+'"]'}function Au(t){return"script[async]"+t}function xd(t,l,e){if(l.count++,l.instance===null)switch(l.type){case"style":var a=t.querySelector('style[data-href~="'+Sl(e.href)+'"]');if(a)return l.instance=a,Xt(a),a;var u=H({},e,{"data-href":e.href,"data-precedence":e.precedence,href:null,precedence:null});return a=(t.ownerDocument||t).createElement("style"),Xt(a),Jt(a,"style",u),Zn(a,e.precedence,t),l.instance=a;case"stylesheet":u=Ca(e.href);var n=t.querySelector(Eu(u));if(n)return l.state.loading|=4,l.instance=n,Xt(n),n;a=Sd(e),(u=Ol.get(u))&&sf(a,u),n=(t.ownerDocument||t).createElement("link"),Xt(n);var i=n;return i._p=new Promise(function(c,o){i.onload=c,i.onerror=o}),Jt(n,"link",a),l.state.loading|=4,Zn(n,e.precedence,t),l.instance=n;case"script":return n=ja(e.src),(u=t.querySelector(Au(n)))?(l.instance=u,Xt(u),u):(a=e,(u=Ol.get(n))&&(a=H({},e),df(a,u)),t=t.ownerDocument||t,u=t.createElement("script"),Xt(u),Jt(u,"link",a),t.head.appendChild(u),l.instance=u);case"void":return null;default:throw Error(f(443,l.type))}else l.type==="stylesheet"&&(l.state.loading&4)===0&&(a=l.instance,l.state.loading|=4,Zn(a,e.precedence,t));return l.instance}function Zn(t,l,e){for(var a=e.querySelectorAll('link[rel="stylesheet"][data-precedence],style[data-precedence]'),u=a.length?a[a.length-1]:null,n=u,i=0;i title"):null)}function xm(t,l,e){if(e===1||l.itemProp!=null)return!1;switch(t){case"meta":case"title":return!0;case"style":if(typeof l.precedence!="string"||typeof l.href!="string"||l.href==="")break;return!0;case"link":if(typeof l.rel!="string"||typeof l.href!="string"||l.href===""||l.onLoad||l.onError)break;return l.rel==="stylesheet"?(t=l.disabled,typeof l.precedence=="string"&&t==null):!0;case"script":if(l.async&&typeof l.async!="function"&&typeof l.async!="symbol"&&!l.onLoad&&!l.onError&&l.src&&typeof l.src=="string")return!0}return!1}function Ed(t){return!(t.type==="stylesheet"&&(t.state.loading&3)===0)}function zm(t,l,e,a){if(e.type==="stylesheet"&&(typeof a.media!="string"||matchMedia(a.media).matches!==!1)&&(e.state.loading&4)===0){if(e.instance===null){var u=Ca(a.href),n=l.querySelector(Eu(u));if(n){l=n._p,l!==null&&typeof l=="object"&&typeof l.then=="function"&&(t.count++,t=Ln.bind(t),l.then(t,t)),e.state.loading|=4,e.instance=n,Xt(n);return}n=l.ownerDocument||l,a=Sd(a),(u=Ol.get(u))&&sf(a,u),n=n.createElement("link"),Xt(n);var i=n;i._p=new Promise(function(c,o){i.onload=c,i.onerror=o}),Jt(n,"link",a),e.instance=n}t.stylesheets===null&&(t.stylesheets=new Map),t.stylesheets.set(e,l),(l=e.state.preload)&&(e.state.loading&3)===0&&(t.count++,e=Ln.bind(t),l.addEventListener("load",e),l.addEventListener("error",e))}}var yf=0;function Tm(t,l){return t.stylesheets&&t.count===0&&Kn(t,t.stylesheets),0yf?50:800)+l);return t.unsuspend=e,function(){t.unsuspend=null,clearTimeout(a),clearTimeout(u)}}:null}function Ln(){if(this.count--,this.count===0&&(this.imgCount===0||!this.waitingForImages)){if(this.stylesheets)Kn(this,this.stylesheets);else if(this.unsuspend){var t=this.unsuspend;this.unsuspend=null,t()}}}var Vn=null;function Kn(t,l){t.stylesheets=null,t.unsuspend!==null&&(t.count++,Vn=new Map,l.forEach(Em,t),Vn=null,Ln.call(t))}function Em(t,l){if(!(l.state.loading&4)){var e=Vn.get(t);if(e)var a=e.get(null);else{e=new Map,Vn.set(t,e);for(var u=t.querySelectorAll("link[data-precedence],style[data-precedence]"),n=0;n"u"||typeof __REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE!="function"))try{__REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE(r)}catch(v){console.error(v)}}return r(),zf.exports=Xm(),zf.exports}var Zm=Qm();const wm=r=>r.replace(/([a-z0-9])([A-Z])/g,"$1-$2").toLowerCase(),Pd=(...r)=>r.filter((v,S,f)=>!!v&&v.trim()!==""&&f.indexOf(v)===S).join(" ").trim();var Lm={xmlns:"http://www.w3.org/2000/svg",width:24,height:24,viewBox:"0 0 24 24",fill:"none",stroke:"currentColor",strokeWidth:2,strokeLinecap:"round",strokeLinejoin:"round"};const Vm=xt.forwardRef(({color:r="currentColor",size:v=24,strokeWidth:S=2,absoluteStrokeWidth:f,className:_="",children:O,iconNode:D,...U},N)=>xt.createElement("svg",{ref:N,...Lm,width:v,height:v,stroke:r,strokeWidth:f?Number(S)*24/Number(v):S,className:Pd("lucide",_),...U},[...D.map(([p,R])=>xt.createElement(p,R)),...Array.isArray(O)?O:[O]]));const Nl=(r,v)=>{const S=xt.forwardRef(({className:f,..._},O)=>xt.createElement(Vm,{ref:O,iconNode:v,className:Pd(`lucide-${wm(r)}`,f),..._}));return S.displayName=`${r}`,S};const Km=Nl("Binary",[["rect",{x:"14",y:"14",width:"4",height:"6",rx:"2",key:"p02svl"}],["rect",{x:"6",y:"4",width:"4",height:"6",rx:"2",key:"xm4xkj"}],["path",{d:"M6 20h4",key:"1i6q5t"}],["path",{d:"M14 10h4",key:"ru81e7"}],["path",{d:"M6 14h2v6",key:"16z9wg"}],["path",{d:"M14 4h2v6",key:"1idq9u"}]]);const Jm=Nl("BookText",[["path",{d:"M4 19.5v-15A2.5 2.5 0 0 1 6.5 2H19a1 1 0 0 1 1 1v18a1 1 0 0 1-1 1H6.5a1 1 0 0 1 0-5H20",key:"k3hazp"}],["path",{d:"M8 11h8",key:"vwpz6n"}],["path",{d:"M8 7h6",key:"1f0q6e"}]]);const km=Nl("EyeOff",[["path",{d:"M10.733 5.076a10.744 10.744 0 0 1 11.205 6.575 1 1 0 0 1 0 .696 10.747 10.747 0 0 1-1.444 2.49",key:"ct8e1f"}],["path",{d:"M14.084 14.158a3 3 0 0 1-4.242-4.242",key:"151rxh"}],["path",{d:"M17.479 17.499a10.75 10.75 0 0 1-15.417-5.151 1 1 0 0 1 0-.696 10.75 10.75 0 0 1 4.446-5.143",key:"13bj9a"}],["path",{d:"m2 2 20 20",key:"1ooewy"}]]);const Wm=Nl("Eye",[["path",{d:"M2.062 12.348a1 1 0 0 1 0-.696 10.75 10.75 0 0 1 19.876 0 1 1 0 0 1 0 .696 10.75 10.75 0 0 1-19.876 0",key:"1nclc0"}],["circle",{cx:"12",cy:"12",r:"3",key:"1v7zrd"}]]);const $m=Nl("Globe",[["circle",{cx:"12",cy:"12",r:"10",key:"1mglay"}],["path",{d:"M12 2a14.5 14.5 0 0 0 0 20 14.5 14.5 0 0 0 0-20",key:"13o1zl"}],["path",{d:"M2 12h20",key:"9i4pu4"}]]);const kd=Nl("LoaderCircle",[["path",{d:"M21 12a9 9 0 1 1-6.219-8.56",key:"13zald"}]]);const Fm=Nl("Lock",[["rect",{width:"18",height:"11",x:"3",y:"11",rx:"2",ry:"2",key:"1w4ew1"}],["path",{d:"M7 11V7a5 5 0 0 1 10 0v4",key:"fwvmzm"}]]);const Im=Nl("LogIn",[["path",{d:"M15 3h4a2 2 0 0 1 2 2v14a2 2 0 0 1-2 2h-4",key:"u53s6r"}],["polyline",{points:"10 17 15 12 10 7",key:"1ail0h"}],["line",{x1:"15",x2:"3",y1:"12",y2:"12",key:"v6grx8"}]]);const Pm=Nl("RotateCw",[["path",{d:"M21 12a9 9 0 1 1-9-9c2.52 0 4.93 1 6.74 2.74L21 8",key:"1p45f6"}],["path",{d:"M21 3v5h-5",key:"1q7to0"}]]);const th=Nl("User",[["path",{d:"M19 21v-2a4 4 0 0 0-4-4H9a4 4 0 0 0-4 4v2",key:"975kel"}],["circle",{cx:"12",cy:"7",r:"4",key:"17ys0d"}]]);const lh=Nl("Waypoints",[["circle",{cx:"12",cy:"4.5",r:"2.5",key:"r5ysbb"}],["path",{d:"m10.2 6.3-3.9 3.9",key:"1nzqf6"}],["circle",{cx:"4.5",cy:"12",r:"2.5",key:"jydg6v"}],["path",{d:"M7 12h10",key:"b7w52i"}],["circle",{cx:"19.5",cy:"12",r:"2.5",key:"1piiel"}],["path",{d:"m13.8 17.7 3.9-3.9",key:"1wyg1y"}],["circle",{cx:"12",cy:"19.5",r:"2.5",key:"13o1pw"}]]);const eh=Nl("X",[["path",{d:"M18 6 6 18",key:"1bl5f8"}],["path",{d:"m6 6 12 12",key:"d8bk6v"}]]);function t0(){return globalThis.__DATA__??{}}function l0(r){var v,S,f="";if(typeof r=="string"||typeof r=="number")f+=r;else if(typeof r=="object")if(Array.isArray(r)){var _=r.length;for(v=0;v<_;v++)r[v]&&(S=l0(r[v]))&&(f&&(f+=" "),f+=S)}else for(S in r)r[S]&&(f&&(f+=" "),f+=S);return f}function ah(){for(var r,v,S=0,f="",_=arguments.length;S<_;S++)(r=arguments[S])&&(v=l0(r))&&(f&&(f+=" "),f+=v);return f}const Hf="-",uh=r=>{const v=ih(r),{conflictingClassGroups:S,conflictingClassGroupModifiers:f}=r;return{getClassGroupId:D=>{const U=D.split(Hf);return U[0]===""&&U.length!==1&&U.shift(),e0(U,v)||nh(D)},getConflictingClassGroupIds:(D,U)=>{const N=S[D]||[];return U&&f[D]?[...N,...f[D]]:N}}},e0=(r,v)=>{if(r.length===0)return v.classGroupId;const S=r[0],f=v.nextPart.get(S),_=f?e0(r.slice(1),f):void 0;if(_)return _;if(v.validators.length===0)return;const O=r.join(Hf);return v.validators.find(({validator:D})=>D(O))?.classGroupId},Wd=/^\[(.+)\]$/,nh=r=>{if(Wd.test(r)){const v=Wd.exec(r)[1],S=v?.substring(0,v.indexOf(":"));if(S)return"arbitrary.."+S}},ih=r=>{const{theme:v,prefix:S}=r,f={nextPart:new Map,validators:[]};return fh(Object.entries(r.classGroups),S).forEach(([O,D])=>{Df(D,f,O,v)}),f},Df=(r,v,S,f)=>{r.forEach(_=>{if(typeof _=="string"){const O=_===""?v:$d(v,_);O.classGroupId=S;return}if(typeof _=="function"){if(ch(_)){Df(_(f),v,S,f);return}v.validators.push({validator:_,classGroupId:S});return}Object.entries(_).forEach(([O,D])=>{Df(D,$d(v,O),S,f)})})},$d=(r,v)=>{let S=r;return v.split(Hf).forEach(f=>{S.nextPart.has(f)||S.nextPart.set(f,{nextPart:new Map,validators:[]}),S=S.nextPart.get(f)}),S},ch=r=>r.isThemeGetter,fh=(r,v)=>v?r.map(([S,f])=>{const _=f.map(O=>typeof O=="string"?v+O:typeof O=="object"?Object.fromEntries(Object.entries(O).map(([D,U])=>[v+D,U])):O);return[S,_]}):r,rh=r=>{if(r<1)return{get:()=>{},set:()=>{}};let v=0,S=new Map,f=new Map;const _=(O,D)=>{S.set(O,D),v++,v>r&&(v=0,f=S,S=new Map)};return{get(O){let D=S.get(O);if(D!==void 0)return D;if((D=f.get(O))!==void 0)return _(O,D),D},set(O,D){S.has(O)?S.set(O,D):_(O,D)}}},a0="!",oh=r=>{const{separator:v,experimentalParseClassName:S}=r,f=v.length===1,_=v[0],O=v.length,D=U=>{const N=[];let p=0,R=0,H;for(let Q=0;QR?H-R:void 0;return{modifiers:N,hasImportantModifier:st,baseClassName:ct,maybePostfixModifierPosition:G}};return S?U=>S({className:U,parseClassName:D}):D},sh=r=>{if(r.length<=1)return r;const v=[];let S=[];return r.forEach(f=>{f[0]==="["?(v.push(...S.sort(),f),S=[]):S.push(f)}),v.push(...S.sort()),v},dh=r=>({cache:rh(r.cacheSize),parseClassName:oh(r),...uh(r)}),yh=/\s+/,mh=(r,v)=>{const{parseClassName:S,getClassGroupId:f,getConflictingClassGroupIds:_}=v,O=[],D=r.trim().split(yh);let U="";for(let N=D.length-1;N>=0;N-=1){const p=D[N],{modifiers:R,hasImportantModifier:H,baseClassName:V,maybePostfixModifierPosition:st}=S(p);let ct=!!st,G=f(ct?V.substring(0,st):V);if(!G){if(!ct){U=p+(U.length>0?" "+U:U);continue}if(G=f(V),!G){U=p+(U.length>0?" "+U:U);continue}ct=!1}const Q=sh(R).join(":"),L=H?Q+a0:Q,gt=L+G;if(O.includes(gt))continue;O.push(gt);const zt=_(G,ct);for(let _t=0;_t0?" "+U:U)}return U};function hh(){let r=0,v,S,f="";for(;r{if(typeof r=="string")return r;let v,S="";for(let f=0;fH(R),r());return S=dh(p),f=S.cache.get,_=S.cache.set,O=U,U(N)}function U(N){const p=f(N);if(p)return p;const R=mh(N,S);return _(N,R),R}return function(){return O(hh.apply(null,arguments))}}const At=r=>{const v=S=>S[r]||[];return v.isThemeGetter=!0,v},n0=/^\[(?:([a-z-]+):)?(.+)\]$/i,vh=/^\d+\/\d+$/,bh=new Set(["px","full","screen"]),ph=/^(\d+(\.\d+)?)?(xs|sm|md|lg|xl)$/,Sh=/\d+(%|px|r?em|[sdl]?v([hwib]|min|max)|pt|pc|in|cm|mm|cap|ch|ex|r?lh|cq(w|h|i|b|min|max))|\b(calc|min|max|clamp)\(.+\)|^0$/,xh=/^(rgba?|hsla?|hwb|(ok)?(lab|lch)|color-mix)\(.+\)$/,zh=/^(inset_)?-?((\d+)?\.?(\d+)[a-z]+|0)_-?((\d+)?\.?(\d+)[a-z]+|0)/,Th=/^(url|image|image-set|cross-fade|element|(repeating-)?(linear|radial|conic)-gradient)\(.+\)$/,ae=r=>Ha(r)||bh.has(r)||vh.test(r),Ne=r=>Ba(r,"length",Uh),Ha=r=>!!r&&!Number.isNaN(Number(r)),Mf=r=>Ba(r,"number",Ha),Cu=r=>!!r&&Number.isInteger(Number(r)),Eh=r=>r.endsWith("%")&&Ha(r.slice(0,-1)),F=r=>n0.test(r),De=r=>ph.test(r),Ah=new Set(["length","size","percentage"]),Mh=r=>Ba(r,Ah,i0),_h=r=>Ba(r,"position",i0),Oh=new Set(["image","url"]),Nh=r=>Ba(r,Oh,jh),Dh=r=>Ba(r,"",Ch),ju=()=>!0,Ba=(r,v,S)=>{const f=n0.exec(r);return f?f[1]?typeof v=="string"?f[1]===v:v.has(f[1]):S(f[2]):!1},Uh=r=>Sh.test(r)&&!xh.test(r),i0=()=>!1,Ch=r=>zh.test(r),jh=r=>Th.test(r),Rh=()=>{const r=At("colors"),v=At("spacing"),S=At("blur"),f=At("brightness"),_=At("borderColor"),O=At("borderRadius"),D=At("borderSpacing"),U=At("borderWidth"),N=At("contrast"),p=At("grayscale"),R=At("hueRotate"),H=At("invert"),V=At("gap"),st=At("gradientColorStops"),ct=At("gradientColorStopPositions"),G=At("inset"),Q=At("margin"),L=At("opacity"),gt=At("padding"),zt=At("saturate"),_t=At("scale"),it=At("sepia"),Ot=At("skew"),J=At("space"),Rt=At("translate"),It=()=>["auto","contain","none"],jl=()=>["auto","hidden","clip","visible","scroll"],Pt=()=>["auto",F,v],I=()=>[F,v],Rl=()=>["",ae,Ne],tl=()=>["auto",Ha,F],ll=()=>["bottom","center","left","left-bottom","left-top","right","right-bottom","right-top","top"],x=()=>["solid","dashed","dotted","double","none"],C=()=>["normal","multiply","screen","overlay","darken","lighten","color-dodge","color-burn","hard-light","soft-light","difference","exclusion","hue","saturation","color","luminosity"],Z=()=>["start","end","center","between","around","evenly","stretch"],nt=()=>["","0",F],dt=()=>["auto","avoid","all","avoid-page","page","left","right","column"],s=()=>[Ha,F];return{cacheSize:500,separator:":",theme:{colors:[ju],spacing:[ae,Ne],blur:["none","",De,F],brightness:s(),borderColor:[r],borderRadius:["none","","full",De,F],borderSpacing:I(),borderWidth:Rl(),contrast:s(),grayscale:nt(),hueRotate:s(),invert:nt(),gap:I(),gradientColorStops:[r],gradientColorStopPositions:[Eh,Ne],inset:Pt(),margin:Pt(),opacity:s(),padding:I(),saturate:s(),scale:s(),sepia:nt(),skew:s(),space:I(),translate:I()},classGroups:{aspect:[{aspect:["auto","square","video",F]}],container:["container"],columns:[{columns:[De]}],"break-after":[{"break-after":dt()}],"break-before":[{"break-before":dt()}],"break-inside":[{"break-inside":["auto","avoid","avoid-page","avoid-column"]}],"box-decoration":[{"box-decoration":["slice","clone"]}],box:[{box:["border","content"]}],display:["block","inline-block","inline","flex","inline-flex","table","inline-table","table-caption","table-cell","table-column","table-column-group","table-footer-group","table-header-group","table-row-group","table-row","flow-root","grid","inline-grid","contents","list-item","hidden"],float:[{float:["right","left","none","start","end"]}],clear:[{clear:["left","right","both","none","start","end"]}],isolation:["isolate","isolation-auto"],"object-fit":[{object:["contain","cover","fill","none","scale-down"]}],"object-position":[{object:[...ll(),F]}],overflow:[{overflow:jl()}],"overflow-x":[{"overflow-x":jl()}],"overflow-y":[{"overflow-y":jl()}],overscroll:[{overscroll:It()}],"overscroll-x":[{"overscroll-x":It()}],"overscroll-y":[{"overscroll-y":It()}],position:["static","fixed","absolute","relative","sticky"],inset:[{inset:[G]}],"inset-x":[{"inset-x":[G]}],"inset-y":[{"inset-y":[G]}],start:[{start:[G]}],end:[{end:[G]}],top:[{top:[G]}],right:[{right:[G]}],bottom:[{bottom:[G]}],left:[{left:[G]}],visibility:["visible","invisible","collapse"],z:[{z:["auto",Cu,F]}],basis:[{basis:Pt()}],"flex-direction":[{flex:["row","row-reverse","col","col-reverse"]}],"flex-wrap":[{flex:["wrap","wrap-reverse","nowrap"]}],flex:[{flex:["1","auto","initial","none",F]}],grow:[{grow:nt()}],shrink:[{shrink:nt()}],order:[{order:["first","last","none",Cu,F]}],"grid-cols":[{"grid-cols":[ju]}],"col-start-end":[{col:["auto",{span:["full",Cu,F]},F]}],"col-start":[{"col-start":tl()}],"col-end":[{"col-end":tl()}],"grid-rows":[{"grid-rows":[ju]}],"row-start-end":[{row:["auto",{span:[Cu,F]},F]}],"row-start":[{"row-start":tl()}],"row-end":[{"row-end":tl()}],"grid-flow":[{"grid-flow":["row","col","dense","row-dense","col-dense"]}],"auto-cols":[{"auto-cols":["auto","min","max","fr",F]}],"auto-rows":[{"auto-rows":["auto","min","max","fr",F]}],gap:[{gap:[V]}],"gap-x":[{"gap-x":[V]}],"gap-y":[{"gap-y":[V]}],"justify-content":[{justify:["normal",...Z()]}],"justify-items":[{"justify-items":["start","end","center","stretch"]}],"justify-self":[{"justify-self":["auto","start","end","center","stretch"]}],"align-content":[{content:["normal",...Z(),"baseline"]}],"align-items":[{items:["start","end","center","baseline","stretch"]}],"align-self":[{self:["auto","start","end","center","stretch","baseline"]}],"place-content":[{"place-content":[...Z(),"baseline"]}],"place-items":[{"place-items":["start","end","center","baseline","stretch"]}],"place-self":[{"place-self":["auto","start","end","center","stretch"]}],p:[{p:[gt]}],px:[{px:[gt]}],py:[{py:[gt]}],ps:[{ps:[gt]}],pe:[{pe:[gt]}],pt:[{pt:[gt]}],pr:[{pr:[gt]}],pb:[{pb:[gt]}],pl:[{pl:[gt]}],m:[{m:[Q]}],mx:[{mx:[Q]}],my:[{my:[Q]}],ms:[{ms:[Q]}],me:[{me:[Q]}],mt:[{mt:[Q]}],mr:[{mr:[Q]}],mb:[{mb:[Q]}],ml:[{ml:[Q]}],"space-x":[{"space-x":[J]}],"space-x-reverse":["space-x-reverse"],"space-y":[{"space-y":[J]}],"space-y-reverse":["space-y-reverse"],w:[{w:["auto","min","max","fit","svw","lvw","dvw",F,v]}],"min-w":[{"min-w":[F,v,"min","max","fit"]}],"max-w":[{"max-w":[F,v,"none","full","min","max","fit","prose",{screen:[De]},De]}],h:[{h:[F,v,"auto","min","max","fit","svh","lvh","dvh"]}],"min-h":[{"min-h":[F,v,"min","max","fit","svh","lvh","dvh"]}],"max-h":[{"max-h":[F,v,"min","max","fit","svh","lvh","dvh"]}],size:[{size:[F,v,"auto","min","max","fit"]}],"font-size":[{text:["base",De,Ne]}],"font-smoothing":["antialiased","subpixel-antialiased"],"font-style":["italic","not-italic"],"font-weight":[{font:["thin","extralight","light","normal","medium","semibold","bold","extrabold","black",Mf]}],"font-family":[{font:[ju]}],"fvn-normal":["normal-nums"],"fvn-ordinal":["ordinal"],"fvn-slashed-zero":["slashed-zero"],"fvn-figure":["lining-nums","oldstyle-nums"],"fvn-spacing":["proportional-nums","tabular-nums"],"fvn-fraction":["diagonal-fractions","stacked-fractions"],tracking:[{tracking:["tighter","tight","normal","wide","wider","widest",F]}],"line-clamp":[{"line-clamp":["none",Ha,Mf]}],leading:[{leading:["none","tight","snug","normal","relaxed","loose",ae,F]}],"list-image":[{"list-image":["none",F]}],"list-style-type":[{list:["none","disc","decimal",F]}],"list-style-position":[{list:["inside","outside"]}],"placeholder-color":[{placeholder:[r]}],"placeholder-opacity":[{"placeholder-opacity":[L]}],"text-alignment":[{text:["left","center","right","justify","start","end"]}],"text-color":[{text:[r]}],"text-opacity":[{"text-opacity":[L]}],"text-decoration":["underline","overline","line-through","no-underline"],"text-decoration-style":[{decoration:[...x(),"wavy"]}],"text-decoration-thickness":[{decoration:["auto","from-font",ae,Ne]}],"underline-offset":[{"underline-offset":["auto",ae,F]}],"text-decoration-color":[{decoration:[r]}],"text-transform":["uppercase","lowercase","capitalize","normal-case"],"text-overflow":["truncate","text-ellipsis","text-clip"],"text-wrap":[{text:["wrap","nowrap","balance","pretty"]}],indent:[{indent:I()}],"vertical-align":[{align:["baseline","top","middle","bottom","text-top","text-bottom","sub","super",F]}],whitespace:[{whitespace:["normal","nowrap","pre","pre-line","pre-wrap","break-spaces"]}],break:[{break:["normal","words","all","keep"]}],hyphens:[{hyphens:["none","manual","auto"]}],content:[{content:["none",F]}],"bg-attachment":[{bg:["fixed","local","scroll"]}],"bg-clip":[{"bg-clip":["border","padding","content","text"]}],"bg-opacity":[{"bg-opacity":[L]}],"bg-origin":[{"bg-origin":["border","padding","content"]}],"bg-position":[{bg:[...ll(),_h]}],"bg-repeat":[{bg:["no-repeat",{repeat:["","x","y","round","space"]}]}],"bg-size":[{bg:["auto","cover","contain",Mh]}],"bg-image":[{bg:["none",{"gradient-to":["t","tr","r","br","b","bl","l","tl"]},Nh]}],"bg-color":[{bg:[r]}],"gradient-from-pos":[{from:[ct]}],"gradient-via-pos":[{via:[ct]}],"gradient-to-pos":[{to:[ct]}],"gradient-from":[{from:[st]}],"gradient-via":[{via:[st]}],"gradient-to":[{to:[st]}],rounded:[{rounded:[O]}],"rounded-s":[{"rounded-s":[O]}],"rounded-e":[{"rounded-e":[O]}],"rounded-t":[{"rounded-t":[O]}],"rounded-r":[{"rounded-r":[O]}],"rounded-b":[{"rounded-b":[O]}],"rounded-l":[{"rounded-l":[O]}],"rounded-ss":[{"rounded-ss":[O]}],"rounded-se":[{"rounded-se":[O]}],"rounded-ee":[{"rounded-ee":[O]}],"rounded-es":[{"rounded-es":[O]}],"rounded-tl":[{"rounded-tl":[O]}],"rounded-tr":[{"rounded-tr":[O]}],"rounded-br":[{"rounded-br":[O]}],"rounded-bl":[{"rounded-bl":[O]}],"border-w":[{border:[U]}],"border-w-x":[{"border-x":[U]}],"border-w-y":[{"border-y":[U]}],"border-w-s":[{"border-s":[U]}],"border-w-e":[{"border-e":[U]}],"border-w-t":[{"border-t":[U]}],"border-w-r":[{"border-r":[U]}],"border-w-b":[{"border-b":[U]}],"border-w-l":[{"border-l":[U]}],"border-opacity":[{"border-opacity":[L]}],"border-style":[{border:[...x(),"hidden"]}],"divide-x":[{"divide-x":[U]}],"divide-x-reverse":["divide-x-reverse"],"divide-y":[{"divide-y":[U]}],"divide-y-reverse":["divide-y-reverse"],"divide-opacity":[{"divide-opacity":[L]}],"divide-style":[{divide:x()}],"border-color":[{border:[_]}],"border-color-x":[{"border-x":[_]}],"border-color-y":[{"border-y":[_]}],"border-color-s":[{"border-s":[_]}],"border-color-e":[{"border-e":[_]}],"border-color-t":[{"border-t":[_]}],"border-color-r":[{"border-r":[_]}],"border-color-b":[{"border-b":[_]}],"border-color-l":[{"border-l":[_]}],"divide-color":[{divide:[_]}],"outline-style":[{outline:["",...x()]}],"outline-offset":[{"outline-offset":[ae,F]}],"outline-w":[{outline:[ae,Ne]}],"outline-color":[{outline:[r]}],"ring-w":[{ring:Rl()}],"ring-w-inset":["ring-inset"],"ring-color":[{ring:[r]}],"ring-opacity":[{"ring-opacity":[L]}],"ring-offset-w":[{"ring-offset":[ae,Ne]}],"ring-offset-color":[{"ring-offset":[r]}],shadow:[{shadow:["","inner","none",De,Dh]}],"shadow-color":[{shadow:[ju]}],opacity:[{opacity:[L]}],"mix-blend":[{"mix-blend":[...C(),"plus-lighter","plus-darker"]}],"bg-blend":[{"bg-blend":C()}],filter:[{filter:["","none"]}],blur:[{blur:[S]}],brightness:[{brightness:[f]}],contrast:[{contrast:[N]}],"drop-shadow":[{"drop-shadow":["","none",De,F]}],grayscale:[{grayscale:[p]}],"hue-rotate":[{"hue-rotate":[R]}],invert:[{invert:[H]}],saturate:[{saturate:[zt]}],sepia:[{sepia:[it]}],"backdrop-filter":[{"backdrop-filter":["","none"]}],"backdrop-blur":[{"backdrop-blur":[S]}],"backdrop-brightness":[{"backdrop-brightness":[f]}],"backdrop-contrast":[{"backdrop-contrast":[N]}],"backdrop-grayscale":[{"backdrop-grayscale":[p]}],"backdrop-hue-rotate":[{"backdrop-hue-rotate":[R]}],"backdrop-invert":[{"backdrop-invert":[H]}],"backdrop-opacity":[{"backdrop-opacity":[L]}],"backdrop-saturate":[{"backdrop-saturate":[zt]}],"backdrop-sepia":[{"backdrop-sepia":[it]}],"border-collapse":[{border:["collapse","separate"]}],"border-spacing":[{"border-spacing":[D]}],"border-spacing-x":[{"border-spacing-x":[D]}],"border-spacing-y":[{"border-spacing-y":[D]}],"table-layout":[{table:["auto","fixed"]}],caption:[{caption:["top","bottom"]}],transition:[{transition:["none","all","","colors","opacity","shadow","transform",F]}],duration:[{duration:s()}],ease:[{ease:["linear","in","out","in-out",F]}],delay:[{delay:s()}],animate:[{animate:["none","spin","ping","pulse","bounce",F]}],transform:[{transform:["","gpu","none"]}],scale:[{scale:[_t]}],"scale-x":[{"scale-x":[_t]}],"scale-y":[{"scale-y":[_t]}],rotate:[{rotate:[Cu,F]}],"translate-x":[{"translate-x":[Rt]}],"translate-y":[{"translate-y":[Rt]}],"skew-x":[{"skew-x":[Ot]}],"skew-y":[{"skew-y":[Ot]}],"transform-origin":[{origin:["center","top","top-right","right","bottom-right","bottom","bottom-left","left","top-left",F]}],accent:[{accent:["auto",r]}],appearance:[{appearance:["none","auto"]}],cursor:[{cursor:["auto","default","pointer","wait","text","move","help","not-allowed","none","context-menu","progress","cell","crosshair","vertical-text","alias","copy","no-drop","grab","grabbing","all-scroll","col-resize","row-resize","n-resize","e-resize","s-resize","w-resize","ne-resize","nw-resize","se-resize","sw-resize","ew-resize","ns-resize","nesw-resize","nwse-resize","zoom-in","zoom-out",F]}],"caret-color":[{caret:[r]}],"pointer-events":[{"pointer-events":["none","auto"]}],resize:[{resize:["none","y","x",""]}],"scroll-behavior":[{scroll:["auto","smooth"]}],"scroll-m":[{"scroll-m":I()}],"scroll-mx":[{"scroll-mx":I()}],"scroll-my":[{"scroll-my":I()}],"scroll-ms":[{"scroll-ms":I()}],"scroll-me":[{"scroll-me":I()}],"scroll-mt":[{"scroll-mt":I()}],"scroll-mr":[{"scroll-mr":I()}],"scroll-mb":[{"scroll-mb":I()}],"scroll-ml":[{"scroll-ml":I()}],"scroll-p":[{"scroll-p":I()}],"scroll-px":[{"scroll-px":I()}],"scroll-py":[{"scroll-py":I()}],"scroll-ps":[{"scroll-ps":I()}],"scroll-pe":[{"scroll-pe":I()}],"scroll-pt":[{"scroll-pt":I()}],"scroll-pr":[{"scroll-pr":I()}],"scroll-pb":[{"scroll-pb":I()}],"scroll-pl":[{"scroll-pl":I()}],"snap-align":[{snap:["start","end","center","align-none"]}],"snap-stop":[{snap:["normal","always"]}],"snap-type":[{snap:["none","x","y","both"]}],"snap-strictness":[{snap:["mandatory","proximity"]}],touch:[{touch:["auto","none","manipulation"]}],"touch-x":[{"touch-pan":["x","left","right"]}],"touch-y":[{"touch-pan":["y","up","down"]}],"touch-pz":["touch-pinch-zoom"],select:[{select:["none","text","all","auto"]}],"will-change":[{"will-change":["auto","scroll","contents","transform",F]}],fill:[{fill:[r,"none"]}],"stroke-w":[{stroke:[ae,Ne,Mf]}],stroke:[{stroke:[r,"none"]}],sr:["sr-only","not-sr-only"],"forced-color-adjust":[{"forced-color-adjust":["auto","none"]}]},conflictingClassGroups:{overflow:["overflow-x","overflow-y"],overscroll:["overscroll-x","overscroll-y"],inset:["inset-x","inset-y","start","end","top","right","bottom","left"],"inset-x":["right","left"],"inset-y":["top","bottom"],flex:["basis","grow","shrink"],gap:["gap-x","gap-y"],p:["px","py","ps","pe","pt","pr","pb","pl"],px:["pr","pl"],py:["pt","pb"],m:["mx","my","ms","me","mt","mr","mb","ml"],mx:["mr","ml"],my:["mt","mb"],size:["w","h"],"font-size":["leading"],"fvn-normal":["fvn-ordinal","fvn-slashed-zero","fvn-figure","fvn-spacing","fvn-fraction"],"fvn-ordinal":["fvn-normal"],"fvn-slashed-zero":["fvn-normal"],"fvn-figure":["fvn-normal"],"fvn-spacing":["fvn-normal"],"fvn-fraction":["fvn-normal"],"line-clamp":["display","overflow"],rounded:["rounded-s","rounded-e","rounded-t","rounded-r","rounded-b","rounded-l","rounded-ss","rounded-se","rounded-ee","rounded-es","rounded-tl","rounded-tr","rounded-br","rounded-bl"],"rounded-s":["rounded-ss","rounded-es"],"rounded-e":["rounded-se","rounded-ee"],"rounded-t":["rounded-tl","rounded-tr"],"rounded-r":["rounded-tr","rounded-br"],"rounded-b":["rounded-br","rounded-bl"],"rounded-l":["rounded-tl","rounded-bl"],"border-spacing":["border-spacing-x","border-spacing-y"],"border-w":["border-w-s","border-w-e","border-w-t","border-w-r","border-w-b","border-w-l"],"border-w-x":["border-w-r","border-w-l"],"border-w-y":["border-w-t","border-w-b"],"border-color":["border-color-s","border-color-e","border-color-t","border-color-r","border-color-b","border-color-l"],"border-color-x":["border-color-r","border-color-l"],"border-color-y":["border-color-t","border-color-b"],"scroll-m":["scroll-mx","scroll-my","scroll-ms","scroll-me","scroll-mt","scroll-mr","scroll-mb","scroll-ml"],"scroll-mx":["scroll-mr","scroll-ml"],"scroll-my":["scroll-mt","scroll-mb"],"scroll-p":["scroll-px","scroll-py","scroll-ps","scroll-pe","scroll-pt","scroll-pr","scroll-pb","scroll-pl"],"scroll-px":["scroll-pr","scroll-pl"],"scroll-py":["scroll-pt","scroll-pb"],touch:["touch-x","touch-y","touch-pz"],"touch-x":["touch"],"touch-y":["touch"],"touch-pz":["touch"]},conflictingClassGroupModifiers:{"font-size":["leading"]}}},Hh=gh(Rh);function Zt(...r){return Hh(ah(r))}const Bh=["relative cursor-pointer","text-sm focus:z-10 focus:ring-2 font-medium focus:outline-none whitespace-nowrap shadow-sm","inline-flex gap-2 items-center justify-center transition-colors focus:ring-offset-1","disabled:opacity-40 disabled:cursor-not-allowed disabled:text-nb-gray-300 ring-offset-neutral-950/50"],qh={default:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:focus:ring-zinc-800/50 dark:bg-nb-gray dark:text-gray-400 dark:border-gray-700/30 dark:hover:text-white dark:hover:bg-zinc-800/50"],primary:["dark:focus:ring-netbird-600/50 dark:ring-offset-neutral-950/50 enabled:dark:bg-netbird disabled:dark:bg-nb-gray-910 dark:text-gray-100 enabled:dark:hover:text-white enabled:dark:hover:bg-netbird-500/80","enabled:bg-netbird enabled:text-white enabled:focus:ring-netbird-400/50 enabled:hover:bg-netbird-500"],secondary:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-920 dark:text-gray-400 dark:border-gray-700/40 dark:hover:text-white dark:hover:bg-nb-gray-910"],secondaryLighter:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900/70 dark:text-gray-400 dark:border-gray-700/70 dark:hover:text-white dark:hover:bg-nb-gray-800/60"],input:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-neutral-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900 dark:text-gray-400 dark:border-nb-gray-700 dark:hover:bg-nb-gray-900/80"],dropdown:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-neutral-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900/40 dark:text-gray-400 dark:border-nb-gray-900 dark:hover:bg-nb-gray-900/50"],dotted:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900 border-dashed","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900/30 dark:text-gray-400 dark:border-gray-500/40 dark:hover:text-white dark:hover:bg-zinc-800/50"],tertiary:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:focus:ring-zinc-800/50 dark:bg-white dark:text-gray-800 dark:border-gray-700/40 dark:hover:bg-neutral-200 disabled:dark:bg-nb-gray-920 disabled:dark:text-nb-gray-300"],white:["focus:ring-white/50 bg-white text-gray-800 border-white outline-none hover:bg-neutral-200 disabled:dark:bg-nb-gray-920 disabled:dark:text-nb-gray-300","disabled:dark:bg-nb-gray-900 disabled:dark:text-nb-gray-300 disabled:dark:border-nb-gray-900"],outline:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:focus:ring-zinc-800/50 dark:bg-transparent dark:text-netbird dark:border-netbird dark:hover:bg-nb-gray-900/30"],"danger-outline":["enabled:dark:focus:ring-red-800/20 enabled:dark:focus:bg-red-950/40 enabled:hover:dark:bg-red-950/50 enabled:dark:hover:border-red-800/50 dark:bg-transparent dark:text-red-500"],"danger-text":["dark:bg-transparent dark:text-red-500 dark:hover:text-red-600 dark:border-transparent !px-0 !shadow-none !py-0 focus:ring-red-500/30 dark:ring-offset-neutral-950/50"],"default-outline":["dark:ring-offset-nb-gray-950/50 dark:focus:ring-nb-gray-500/20","dark:bg-transparent dark:text-nb-gray-400 dark:border-transparent dark:hover:text-white dark:hover:bg-nb-gray-900/30 dark:hover:border-nb-gray-800/50","data-[state=open]:dark:text-white data-[state=open]:dark:bg-nb-gray-900/30 data-[state=open]:dark:border-nb-gray-800/50"],danger:["dark:focus:ring-red-700/20 dark:focus:bg-red-700 hover:dark:bg-red-700 dark:hover:border-red-800/50 dark:bg-red-600 dark:text-red-100"]},Yh={xs:"text-xs py-2 px-4",xs2:"text-[0.78rem] py-2 px-4",sm:"text-sm py-2.5 px-4",md:"text-sm py-2.5 px-4",lg:"text-base py-2.5 px-4"},Gh={0:"border",1:"border border-transparent",2:"border border-t-0 border-b-0"},Ru=xt.forwardRef(({variant:r="default",rounded:v=!0,border:S=1,size:f="md",stopPropagation:_=!0,className:O,onClick:D,children:U,...N},p)=>A.jsx("button",{type:"button",...N,ref:p,className:Zt(Bh,qh[r],Yh[f],Gh[S?1:0],v&&"rounded-md",O),onClick:R=>{_&&R.stopPropagation(),D?.(R)},children:U}));Ru.displayName="Button";const Xh={default:["bg-nb-gray-900 placeholder:text-neutral-400/70 border-nb-gray-700","ring-offset-neutral-950/50 focus-visible:ring-neutral-500/20"],darker:["bg-nb-gray-920 placeholder:text-neutral-400/70 border-nb-gray-800","ring-offset-neutral-950/50 focus-visible:ring-neutral-500/20"],error:["bg-nb-gray-900 placeholder:text-neutral-400/70 border-red-500 text-red-500","ring-offset-red-500/10 focus-visible:ring-red-500/10"]},Qh={default:"bg-nb-gray-900 border-nb-gray-700 text-nb-gray-300",error:"bg-nb-gray-900 border-red-500 text-nb-gray-300 text-red-500"},c0=xt.forwardRef(({className:r,type:v,customSuffix:S,customPrefix:f,icon:_,maxWidthClass:O="",error:D,variant:U="default",prefixClassName:N,showPasswordToggle:p=!1,...R},H)=>{const[V,st]=xt.useState(!1),ct=v==="password",G=ct&&V?"text":v,L=(ct&&p?A.jsx("button",{type:"button",onClick:()=>st(!V),className:"hover:text-white transition-all","aria-label":"Toggle password visibility",children:V?A.jsx(km,{size:18}):A.jsx(Wm,{size:18})}):null)||S,gt=D?"error":U;return A.jsxs(A.Fragment,{children:[A.jsxs("div",{className:Zt("flex relative h-[42px]",O),children:[f&&A.jsx("div",{className:Zt(Qh[D?"error":"default"],"flex h-[42px] w-auto rounded-l-md px-3 py-2 text-sm","border items-center whitespace-nowrap",R.disabled&&"opacity-40",N),children:f}),A.jsx("div",{className:Zt("absolute left-0 top-0 h-full flex items-center text-xs text-nb-gray-300 pl-3 leading-[0]",R.disabled&&"opacity-40"),children:_}),A.jsx("input",{type:G,ref:H,...R,className:Zt(Xh[gt],"flex h-[42px] w-full rounded-md px-3 py-2 text-sm","file:bg-transparent file:text-sm file:font-medium file:border-0","focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-offset-2","disabled:cursor-not-allowed disabled:opacity-40","border",f&&"!border-l-0 !rounded-l-none",L&&"!pr-16",_&&"!pl-10",r)}),A.jsx("div",{className:Zt("absolute right-0 top-0 h-full flex items-center text-xs text-nb-gray-300 pr-4 leading-[0] select-none",R.disabled&&"opacity-30"),children:L})]}),D&&A.jsx("p",{className:"text-xs text-red-500 mt-2",children:D})]})});c0.displayName="Input";const Zh=xt.forwardRef(function({value:v,onChange:S,length:f=6,disabled:_=!1,className:O,autoFocus:D=!1},U){const N=xt.useRef([]);xt.useImperativeHandle(U,()=>({focus:()=>{N.current[0]?.focus()}}));const p=v.split("").concat(new Array(f).fill("")).slice(0,f),R=Array.from({length:f},(G,Q)=>`pin-${Q}`),H=(G,Q)=>{if(!/^\d*$/.test(Q))return;const L=[...p];L[G]=Q.slice(-1);const gt=L.join("").replaceAll(/\s/g,"");S(gt),Q&&G{Q.key==="Backspace"&&!p[G]&&G>0&&N.current[G-1]?.focus(),Q.key==="ArrowLeft"&&G>0&&N.current[G-1]?.focus(),Q.key==="ArrowRight"&&G{G.preventDefault();const Q=G.clipboardData.getData("text").replaceAll(/\D/g,"").slice(0,f);S(Q);const L=Math.min(Q.length,f-1);N.current[L]?.focus()},ct=G=>{G.target.select()};return A.jsx("div",{className:Zt("flex gap-2 w-full min-w-0",O),children:p.map((G,Q)=>A.jsx("input",{id:R[Q],ref:L=>{N.current[Q]=L},type:"text",inputMode:"numeric",maxLength:1,value:G,onChange:L=>H(Q,L.target.value),onKeyDown:L=>V(Q,L),onPaste:st,onFocus:ct,disabled:_,autoFocus:D&&Q===0,className:Zt("flex-1 min-w-0 h-[42px] text-center text-sm rounded-md","dark:bg-nb-gray-900 border dark:border-nb-gray-700","dark:placeholder:text-neutral-400/70","focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-offset-2","ring-offset-neutral-200/20 dark:ring-offset-neutral-950/50 dark:focus-visible:ring-neutral-500/20","disabled:cursor-not-allowed disabled:opacity-40")},R[Q]))})}),f0=xt.createContext({value:"",onChange:()=>{}}),r0=()=>xt.useContext(f0);function $e({value:r,defaultValue:v,onChange:S,children:f}){const[_,O]=xt.useState(v??""),D=r??_,U=xt.useCallback(p=>{r===void 0&&O(p),S?.(p)},[r,S]),N=xt.useMemo(()=>({value:D,onChange:U}),[D,U]);return A.jsx(f0.Provider,{value:N,children:A.jsx("div",{children:typeof f=="function"?f({value:D,onChange:U}):f})})}function wh({children:r,className:v}){return A.jsx("div",{role:"tablist",className:Zt("bg-nb-gray-930/70 p-1.5 flex justify-center gap-1 border-nb-gray-900",v),children:r})}function Lh({children:r,value:v,disabled:S=!1,className:f,selected:_,onClick:O}){const D=r0(),U=_??v===D.value;let N="";U?N="bg-nb-gray-900 text-white":S||(N="text-nb-gray-400 hover:bg-nb-gray-900/50");const p=()=>{D.onChange(v),O?.()};return A.jsx("button",{role:"tab",type:"button",disabled:S,"aria-selected":U,onClick:p,className:Zt("px-4 py-2 text-sm rounded-md w-full transition-all cursor-pointer",S&&"opacity-30 cursor-not-allowed",N,f),children:A.jsx("div",{className:"flex items-center w-full justify-center gap-2",children:r})})}function Vh({children:r,value:v,className:S,visible:f}){const _=r0();return f??v===_.value?A.jsx("div",{role:"tabpanel",className:Zt("bg-nb-gray-930/70 px-4 pt-4 pb-5 rounded-b-md border border-t-0 border-nb-gray-900",S),children:r}):null}$e.List=wh;$e.Trigger=Lh;$e.Content=Vh;const Kh="/__netbird__/assets/netbird-full.svg",Jh="data:image/svg+xml,%3csvg%20width='31'%20height='23'%20viewBox='0%200%2031%2023'%20fill='none'%20xmlns='http://www.w3.org/2000/svg'%3e%3cpath%20d='M21.4631%200.523438C17.8173%200.857913%2016.0028%202.95675%2015.3171%204.01871L4.66406%2022.4734H17.5163L30.1929%200.523438H21.4631Z'%20fill='%23F68330'/%3e%3cpath%20d='M17.5265%2022.4737L0%203.88525C0%203.88525%2019.8177%20-1.44128%2021.7493%2015.1738L17.5265%2022.4737Z'%20fill='%23F68330'/%3e%3cpath%20d='M14.9236%204.70563L9.54688%2014.0208L17.5158%2022.4747L21.7385%2015.158C21.0696%209.44682%2018.2851%206.32784%2014.9236%204.69727'%20fill='%23F05252'/%3e%3c/svg%3e",ti={small:{desktop:14,mobile:20},default:{desktop:22,mobile:30},large:{desktop:24,mobile:40}},kh=({size:r="default",mobile:v=!0})=>A.jsxs(A.Fragment,{children:[A.jsx("img",{src:Kh,height:ti[r].desktop,style:{height:ti[r].desktop},alt:"NetBird Logo",className:Zt(v&&"hidden md:block","group-hover:opacity-80 transition-all")}),v&&A.jsx("img",{src:Jh,width:ti[r].mobile,style:{width:ti[r].mobile},alt:"NetBird Logo",className:Zt(v&&"md:hidden ml-4")})]});function Uf(){return A.jsxs("a",{href:"https://netbird.io?utm_source=netbird-proxy&utm_medium=web&utm_campaign=powered_by",target:"_blank",rel:"noopener noreferrer",className:"flex items-center justify-center mt-8 gap-2 group cursor-pointer",children:[A.jsx("span",{className:"text-sm text-nb-gray-400 font-light text-center group-hover:opacity-80 transition-all",children:"Powered by"}),A.jsx(kh,{size:"small",mobile:!1})]})}const Wh=({className:r})=>A.jsx("div",{className:Zt("h-full w-full absolute left-0 top-0 rounded-md overflow-hidden z-0 pointer-events-none",r),children:A.jsx("div",{className:"bg-linear-to-b from-nb-gray-900/10 via-transparent to-transparent w-full h-full rounded-md"})}),Fd=({children:r,className:v})=>A.jsxs("div",{className:Zt("px-6 sm:px-10 py-10 pt-8","bg-nb-gray-940 border border-nb-gray-910 rounded-lg relative",v),children:[A.jsx(Wh,{}),r]});function Cf({children:r,className:v}){return A.jsx("h1",{className:Zt("text-xl! text-center z-10 relative",v),children:r})}function jf({children:r,className:v}){return A.jsx("div",{className:Zt("text-sm text-nb-gray-300 font-light mt-2 block text-center z-10 relative",v),children:r})}const $h=()=>A.jsxs("div",{className:"flex items-center justify-center relative my-4",children:[A.jsx("span",{className:"bg-nb-gray-940 relative z-10 px-4 text-xs text-nb-gray-400 font-medium",children:"OR"}),A.jsx("span",{className:"h-px bg-nb-gray-900 w-full absolute z-0"})]}),Fh=({error:r})=>A.jsx("div",{className:"text-red-400 bg-red-800/20 border border-red-800/50 rounded-lg px-4 py-3 whitespace-break-spaces text-sm",children:r});function Id({className:r,htmlFor:v,...S}){return A.jsx("label",{htmlFor:v,className:Zt("text-sm font-medium tracking-wider leading-none","peer-disabled:cursor-not-allowed peer-disabled:opacity-70","mb-2.5 inline-block text-nb-gray-200","flex items-center gap-2 select-none",r),...S})}const _f=t0(),Ft=_f.methods&&Object.keys(_f.methods).length>0?_f.methods:{password:"password",pin:"pin",oidc:"/auth/oidc"};function Ih(){xt.useEffect(()=>{document.title="Authentication Required - NetBird Service"},[]);const[r,v]=xt.useState(null),[S,f]=xt.useState(null),[_,O]=xt.useState(""),[D,U]=xt.useState(""),N=xt.useRef(null),p=xt.useRef(null),[R,H]=xt.useState(Ft.password?"password":"pin"),V=(it,Ot)=>{v(Ot),f(null),it==="password"?(U(""),setTimeout(()=>N.current?.focus(),200)):(O(""),setTimeout(()=>p.current?.focus(),200))},st=(it,Ot)=>{v(null),f(it);const J=new FormData;it==="password"?J.append(Ft.password,Ot):J.append(Ft.pin,Ot),fetch(globalThis.location.href,{method:"POST",body:J,redirect:"manual"}).then(Rt=>{Rt.type==="opaqueredirect"||Rt.status===0?(f("redirect"),globalThis.location.reload()):V(it,"Authentication failed. Please try again.")}).catch(()=>{V(it,"An error occurred. Please try again.")})},ct=it=>{O(it),it.length===6&&st("pin",it)},G=_.length===6,Q=D.length>0,L=S!==null||R==="password"&&!Q||R==="pin"&&!G,gt=Ft.password||Ft.pin,zt=Ft.password&&Ft.pin,_t=R==="password"?"Sign in":"Submit";return S==="redirect"?A.jsxs("main",{className:"mt-20",children:[A.jsxs(Fd,{className:"max-w-105 mx-auto",children:[A.jsx(Cf,{children:"Authenticated"}),A.jsx(jf,{children:"Loading service..."}),A.jsx("div",{className:"flex justify-center mt-7",children:A.jsx(kd,{className:"animate-spin",size:24})})]}),A.jsx(Uf,{})]}):A.jsxs("main",{className:"mt-20",children:[A.jsxs(Fd,{className:"max-w-105 mx-auto",children:[A.jsx(Cf,{children:"Authentication Required"}),A.jsx(jf,{children:"The service you are trying to access is protected. Please authenticate to continue."}),A.jsxs("div",{className:"flex flex-col gap-4 mt-7 z-10 relative",children:[r&&A.jsx(Fh,{error:r}),Ft.oidc&&A.jsxs(Ru,{variant:"primary",className:"w-full",onClick:()=>{globalThis.location.href=Ft.oidc},children:[A.jsx(Im,{size:16}),"Sign in with SSO"]}),Ft.oidc&>&&A.jsx($h,{}),gt&&A.jsxs("form",{onSubmit:it=>{it.preventDefault(),st(R,R==="password"?D:_)},children:[zt&&A.jsx($e,{value:R,onChange:it=>{H(it),setTimeout(()=>{it==="password"?N.current?.focus():p.current?.focus()},0)},children:A.jsxs($e.List,{className:"rounded-lg border mb-4",children:[A.jsxs($e.Trigger,{value:"password",children:[A.jsx(Fm,{size:14}),"Password"]}),A.jsxs($e.Trigger,{value:"pin",children:[A.jsx(Km,{size:14}),"PIN"]})]})}),A.jsxs("div",{className:"mb-4",children:[Ft.password&&(R==="password"||!Ft.pin)&&A.jsxs(A.Fragment,{children:[!zt&&A.jsx(Id,{htmlFor:"password",children:"Password"}),A.jsx(c0,{ref:N,type:"password",id:"password",placeholder:"Enter password",disabled:S!==null,showPasswordToggle:!0,autoFocus:!0,value:D,onChange:it=>U(it.target.value)})]}),Ft.pin&&(R==="pin"||!Ft.password)&&A.jsxs(A.Fragment,{children:[!zt&&A.jsx(Id,{htmlFor:"pin-0",children:"Enter PIN Code"}),A.jsx(Zh,{ref:p,value:_,onChange:ct,disabled:S!==null,autoFocus:!Ft.password})]})]}),A.jsx(Ru,{type:"submit",disabled:L,variant:"secondary",className:"w-full",children:S===null?_t:A.jsxs(A.Fragment,{children:[A.jsx(kd,{className:"animate-spin",size:16}),"Verifying..."]})})]})]})]}),A.jsx(Uf,{})]})}function Ph({success:r=!0}){return r?A.jsx("div",{className:"flex-1 flex items-center justify-center h-12 w-full px-5",children:A.jsx("div",{className:"w-full border-t-2 border-dashed border-green-500"})}):A.jsxs("div",{className:"flex-1 flex items-center justify-center h-12 min-w-10 px-5 relative",children:[A.jsx("div",{className:"w-full border-t-2 border-dashed border-nb-gray-900"}),A.jsx("div",{className:"absolute inset-0 flex items-center justify-center",children:A.jsx("div",{className:"w-8 h-8 rounded-full flex items-center justify-center",children:A.jsx(eh,{size:18,className:"text-netbird"})})})]})}function Of({icon:r,label:v,detail:S,success:f=!0,line:_=!0}){return A.jsxs(A.Fragment,{children:[_&&A.jsx(Ph,{success:f}),A.jsxs("div",{className:"flex flex-col items-center gap-2",children:[A.jsx("div",{className:"w-14 h-14 rounded-md flex items-center justify-center from-nb-gray-940 to-nb-gray-930/70 bg-gradient-to-br border border-nb-gray-910",children:A.jsx(r,{size:20,className:"text-nb-gray-200"})}),A.jsx("span",{className:"text-sm text-nb-gray-200 font-normal mt-1",children:v}),A.jsx("span",{className:`text-xs font-medium uppercase ${f?"text-green-500":"text-netbird"}`,children:f?"Connected":"Unreachable"}),S&&A.jsx("span",{className:"text-xs text-nb-gray-400 truncate text-center",children:S})]})]})}function tg({code:r,title:v,message:S,proxy:f=!0,destination:_=!0,requestId:O,simple:D=!1,retryUrl:U}){xt.useEffect(()=>{document.title=`${v} - NetBird Service`},[v]);const[N]=xt.useState(()=>new Date().toISOString());return A.jsxs("main",{className:"flex flex-col items-center mt-24 px-4 max-w-3xl mx-auto",children:[A.jsxs("div",{className:"text-sm text-netbird font-normal font-mono mb-3 z-10 relative",children:["Error ",r]}),A.jsx(Cf,{className:"text-3xl!",children:v}),A.jsx(jf,{className:"mt-2 mb-8 max-w-md",children:S}),!D&&A.jsxs("div",{className:"hidden sm:flex items-start justify-center w-full mt-6 mb-16 z-10 relative",children:[A.jsx(Of,{icon:th,label:"You",line:!1}),A.jsx(Of,{icon:lh,label:"Proxy",success:f}),A.jsx(Of,{icon:$m,label:"Destination",success:_})]}),A.jsxs("div",{className:"flex gap-3 justify-center items-center mb-6 z-10 relative",children:[A.jsxs(Ru,{variant:"primary",onClick:()=>{U?globalThis.location.href=U:globalThis.location.reload()},children:[A.jsx(Pm,{size:16}),"Refresh Page"]}),A.jsxs(Ru,{variant:"secondary",onClick:()=>globalThis.open("https://docs.netbird.io","_blank","noopener,noreferrer"),children:[A.jsx(Jm,{size:16}),"Documentation"]})]}),A.jsxs("div",{className:"text-center text-xs text-nb-gray-300 uppercase z-10 relative font-mono flex flex-col sm:flex-row gap-2 sm:gap-10 mt-4 mb-3",children:[A.jsxs("div",{children:[A.jsx("span",{className:"text-nb-gray-400",children:"REQUEST-ID:"})," ",O]}),A.jsxs("div",{children:[A.jsx("span",{className:"text-nb-gray-400",children:"TIMESTAMP:"})," ",N]})]}),A.jsx(Uf,{})]})}const Nf=t0();Zm.createRoot(document.getElementById("root")).render(A.jsx(xt.StrictMode,{children:Nf.page==="error"&&Nf.error?A.jsx(tg,{...Nf.error}):A.jsx(Ih,{})})); +`+a.stack}}var ui=Object.prototype.hasOwnProperty,ni=r.unstable_scheduleCallback,ii=r.unstable_cancelCallback,o0=r.unstable_shouldYield,d0=r.unstable_requestPaint,rl=r.unstable_now,y0=r.unstable_getCurrentPriorityLevel,Yf=r.unstable_ImmediatePriority,Gf=r.unstable_UserBlockingPriority,Bu=r.unstable_NormalPriority,m0=r.unstable_LowPriority,Xf=r.unstable_IdlePriority,h0=r.log,g0=r.unstable_setDisableYieldValue,Ya=null,sl=null;function ue(t){if(typeof h0=="function"&&g0(t),sl&&typeof sl.setStrictMode=="function")try{sl.setStrictMode(Ya,t)}catch{}}var ol=Math.clz32?Math.clz32:p0,v0=Math.log,b0=Math.LN2;function p0(t){return t>>>=0,t===0?32:31-(v0(t)/b0|0)|0}var qu=256,Yu=262144,Gu=4194304;function Ce(t){var l=t&42;if(l!==0)return l;switch(t&-t){case 1:return 1;case 2:return 2;case 4:return 4;case 8:return 8;case 16:return 16;case 32:return 32;case 64:return 64;case 128:return 128;case 256:case 512:case 1024:case 2048:case 4096:case 8192:case 16384:case 32768:case 65536:case 131072:return t&261888;case 262144:case 524288:case 1048576:case 2097152:return t&3932160;case 4194304:case 8388608:case 16777216:case 33554432:return t&62914560;case 67108864:return 67108864;case 134217728:return 134217728;case 268435456:return 268435456;case 536870912:return 536870912;case 1073741824:return 0;default:return t}}function Xu(t,l,e){var a=t.pendingLanes;if(a===0)return 0;var u=0,n=t.suspendedLanes,i=t.pingedLanes;t=t.warmLanes;var c=a&134217727;return c!==0?(a=c&~n,a!==0?u=Ce(a):(i&=c,i!==0?u=Ce(i):e||(e=c&~t,e!==0&&(u=Ce(e))))):(c=a&~n,c!==0?u=Ce(c):i!==0?u=Ce(i):e||(e=a&~t,e!==0&&(u=Ce(e)))),u===0?0:l!==0&&l!==u&&(l&n)===0&&(n=u&-u,e=l&-l,n>=e||n===32&&(e&4194048)!==0)?l:u}function Ga(t,l){return(t.pendingLanes&~(t.suspendedLanes&~t.pingedLanes)&l)===0}function S0(t,l){switch(t){case 1:case 2:case 4:case 8:case 64:return l+250;case 16:case 32:case 128:case 256:case 512:case 1024:case 2048:case 4096:case 8192:case 16384:case 32768:case 65536:case 131072:case 262144:case 524288:case 1048576:case 2097152:return l+5e3;case 4194304:case 8388608:case 16777216:case 33554432:return-1;case 67108864:case 134217728:case 268435456:case 536870912:case 1073741824:return-1;default:return-1}}function Qf(){var t=Gu;return Gu<<=1,(Gu&62914560)===0&&(Gu=4194304),t}function ci(t){for(var l=[],e=0;31>e;e++)l.push(t);return l}function Xa(t,l){t.pendingLanes|=l,l!==268435456&&(t.suspendedLanes=0,t.pingedLanes=0,t.warmLanes=0)}function x0(t,l,e,a,u,n){var i=t.pendingLanes;t.pendingLanes=e,t.suspendedLanes=0,t.pingedLanes=0,t.warmLanes=0,t.expiredLanes&=e,t.entangledLanes&=e,t.errorRecoveryDisabledLanes&=e,t.shellSuspendCounter=0;var c=t.entanglements,s=t.expirationTimes,h=t.hiddenUpdates;for(e=i&~e;0"u")return null;try{return t.activeElement||t.body}catch{return t.body}}var _0=/[\n"\\]/g;function xl(t){return t.replace(_0,function(l){return"\\"+l.charCodeAt(0).toString(16)+" "})}function yi(t,l,e,a,u,n,i,c){t.name="",i!=null&&typeof i!="function"&&typeof i!="symbol"&&typeof i!="boolean"?t.type=i:t.removeAttribute("type"),l!=null?i==="number"?(l===0&&t.value===""||t.value!=l)&&(t.value=""+Sl(l)):t.value!==""+Sl(l)&&(t.value=""+Sl(l)):i!=="submit"&&i!=="reset"||t.removeAttribute("value"),l!=null?mi(t,i,Sl(l)):e!=null?mi(t,i,Sl(e)):a!=null&&t.removeAttribute("value"),u==null&&n!=null&&(t.defaultChecked=!!n),u!=null&&(t.checked=u&&typeof u!="function"&&typeof u!="symbol"),c!=null&&typeof c!="function"&&typeof c!="symbol"&&typeof c!="boolean"?t.name=""+Sl(c):t.removeAttribute("name")}function tr(t,l,e,a,u,n,i,c){if(n!=null&&typeof n!="function"&&typeof n!="symbol"&&typeof n!="boolean"&&(t.type=n),l!=null||e!=null){if(!(n!=="submit"&&n!=="reset"||l!=null)){di(t);return}e=e!=null?""+Sl(e):"",l=l!=null?""+Sl(l):e,c||l===t.value||(t.value=l),t.defaultValue=l}a=a??u,a=typeof a!="function"&&typeof a!="symbol"&&!!a,t.checked=c?t.checked:!!a,t.defaultChecked=!!a,i!=null&&typeof i!="function"&&typeof i!="symbol"&&typeof i!="boolean"&&(t.name=i),di(t)}function mi(t,l,e){l==="number"&&wu(t.ownerDocument)===t||t.defaultValue===""+e||(t.defaultValue=""+e)}function ea(t,l,e,a){if(t=t.options,l){l={};for(var u=0;u"u"||typeof window.document>"u"||typeof window.document.createElement>"u"),pi=!1;if(Ql)try{var La={};Object.defineProperty(La,"passive",{get:function(){pi=!0}}),window.addEventListener("test",La,La),window.removeEventListener("test",La,La)}catch{pi=!1}var ie=null,Si=null,Vu=null;function cr(){if(Vu)return Vu;var t,l=Si,e=l.length,a,u="value"in ie?ie.value:ie.textContent,n=u.length;for(t=0;t=Ja),yr=" ",mr=!1;function hr(t,l){switch(t){case"keyup":return ly.indexOf(l.keyCode)!==-1;case"keydown":return l.keyCode!==229;case"keypress":case"mousedown":case"focusout":return!0;default:return!1}}function gr(t){return t=t.detail,typeof t=="object"&&"data"in t?t.data:null}var ia=!1;function ay(t,l){switch(t){case"compositionend":return gr(l);case"keypress":return l.which!==32?null:(mr=!0,yr);case"textInput":return t=l.data,t===yr&&mr?null:t;default:return null}}function uy(t,l){if(ia)return t==="compositionend"||!Ei&&hr(t,l)?(t=cr(),Vu=Si=ie=null,ia=!1,t):null;switch(t){case"paste":return null;case"keypress":if(!(l.ctrlKey||l.altKey||l.metaKey)||l.ctrlKey&&l.altKey){if(l.char&&1=l)return{node:e,offset:l-t};t=a}t:{for(;e;){if(e.nextSibling){e=e.nextSibling;break t}e=e.parentNode}e=void 0}e=Ar(e)}}function Mr(t,l){return t&&l?t===l?!0:t&&t.nodeType===3?!1:l&&l.nodeType===3?Mr(t,l.parentNode):"contains"in t?t.contains(l):t.compareDocumentPosition?!!(t.compareDocumentPosition(l)&16):!1:!1}function _r(t){t=t!=null&&t.ownerDocument!=null&&t.ownerDocument.defaultView!=null?t.ownerDocument.defaultView:window;for(var l=wu(t.document);l instanceof t.HTMLIFrameElement;){try{var e=typeof l.contentWindow.location.href=="string"}catch{e=!1}if(e)t=l.contentWindow;else break;l=wu(t.document)}return l}function Oi(t){var l=t&&t.nodeName&&t.nodeName.toLowerCase();return l&&(l==="input"&&(t.type==="text"||t.type==="search"||t.type==="tel"||t.type==="url"||t.type==="password")||l==="textarea"||t.contentEditable==="true")}var dy=Ql&&"documentMode"in document&&11>=document.documentMode,ca=null,Ni=null,Fa=null,Di=!1;function Or(t,l,e){var a=e.window===e?e.document:e.nodeType===9?e:e.ownerDocument;Di||ca==null||ca!==wu(a)||(a=ca,"selectionStart"in a&&Oi(a)?a={start:a.selectionStart,end:a.selectionEnd}:(a=(a.ownerDocument&&a.ownerDocument.defaultView||window).getSelection(),a={anchorNode:a.anchorNode,anchorOffset:a.anchorOffset,focusNode:a.focusNode,focusOffset:a.focusOffset}),Fa&&$a(Fa,a)||(Fa=a,a=Gn(Ni,"onSelect"),0>=i,u-=i,Hl=1<<32-ol(l)+u|e<$?(at=Y,Y=null):at=Y.sibling;var rt=g(y,Y,m[$],T);if(rt===null){Y===null&&(Y=at);break}t&&Y&&rt.alternate===null&&l(y,Y),d=n(rt,d,$),ft===null?X=rt:ft.sibling=rt,ft=rt,Y=at}if($===m.length)return e(y,Y),ut&&wl(y,$),X;if(Y===null){for(;$$?(at=Y,Y=null):at=Y.sibling;var Oe=g(y,Y,rt.value,T);if(Oe===null){Y===null&&(Y=at);break}t&&Y&&Oe.alternate===null&&l(y,Y),d=n(Oe,d,$),ft===null?X=Oe:ft.sibling=Oe,ft=Oe,Y=at}if(rt.done)return e(y,Y),ut&&wl(y,$),X;if(Y===null){for(;!rt.done;$++,rt=m.next())rt=A(y,rt.value,T),rt!==null&&(d=n(rt,d,$),ft===null?X=rt:ft.sibling=rt,ft=rt);return ut&&wl(y,$),X}for(Y=a(Y);!rt.done;$++,rt=m.next())rt=b(Y,y,$,rt.value,T),rt!==null&&(t&&rt.alternate!==null&&Y.delete(rt.key===null?$:rt.key),d=n(rt,d,$),ft===null?X=rt:ft.sibling=rt,ft=rt);return t&&Y.forEach(function(Cm){return l(y,Cm)}),ut&&wl(y,$),X}function pt(y,d,m,T){if(typeof m=="object"&&m!==null&&m.type===G&&m.key===null&&(m=m.props.children),typeof m=="object"&&m!==null){switch(m.$$typeof){case ot:t:{for(var X=m.key;d!==null;){if(d.key===X){if(X=m.type,X===G){if(d.tag===7){e(y,d.sibling),T=u(d,m.props.children),T.return=y,y=T;break t}}else if(d.elementType===X||typeof X=="object"&&X!==null&&X.$$typeof===Nt&&we(X)===d.type){e(y,d.sibling),T=u(d,m.props),au(T,m),T.return=y,y=T;break t}e(y,d);break}else l(y,d);d=d.sibling}m.type===G?(T=Ye(m.props.children,y.mode,T,m.key),T.return=y,y=T):(T=ln(m.type,m.key,m.props,null,y.mode,T),au(T,m),T.return=y,y=T)}return i(y);case ct:t:{for(X=m.key;d!==null;){if(d.key===X)if(d.tag===4&&d.stateNode.containerInfo===m.containerInfo&&d.stateNode.implementation===m.implementation){e(y,d.sibling),T=u(d,m.children||[]),T.return=y,y=T;break t}else{e(y,d);break}else l(y,d);d=d.sibling}T=qi(m,y.mode,T),T.return=y,y=T}return i(y);case Nt:return m=we(m),pt(y,d,m,T)}if(ll(m))return B(y,d,m,T);if(I(m)){if(X=I(m),typeof X!="function")throw Error(f(150));return m=X.call(m),w(y,d,m,T)}if(typeof m.then=="function")return pt(y,d,rn(m),T);if(m.$$typeof===zt)return pt(y,d,un(y,m),T);sn(y,m)}return typeof m=="string"&&m!==""||typeof m=="number"||typeof m=="bigint"?(m=""+m,d!==null&&d.tag===6?(e(y,d.sibling),T=u(d,m),T.return=y,y=T):(e(y,d),T=Bi(m,y.mode,T),T.return=y,y=T),i(y)):e(y,d)}return function(y,d,m,T){try{eu=0;var X=pt(y,d,m,T);return ba=null,X}catch(Y){if(Y===va||Y===cn)throw Y;var ft=yl(29,Y,null,y.mode);return ft.lanes=T,ft.return=y,ft}}}var Ve=Fr(!0),Ir=Fr(!1),oe=!1;function Wi(t){t.updateQueue={baseState:t.memoizedState,firstBaseUpdate:null,lastBaseUpdate:null,shared:{pending:null,lanes:0,hiddenCallbacks:null},callbacks:null}}function $i(t,l){t=t.updateQueue,l.updateQueue===t&&(l.updateQueue={baseState:t.baseState,firstBaseUpdate:t.firstBaseUpdate,lastBaseUpdate:t.lastBaseUpdate,shared:t.shared,callbacks:null})}function de(t){return{lane:t,tag:0,payload:null,callback:null,next:null}}function ye(t,l,e){var a=t.updateQueue;if(a===null)return null;if(a=a.shared,(st&2)!==0){var u=a.pending;return u===null?l.next=l:(l.next=u.next,u.next=l),a.pending=l,l=tn(t),Hr(t,null,e),l}return Pu(t,a,l,e),tn(t)}function uu(t,l,e){if(l=l.updateQueue,l!==null&&(l=l.shared,(e&4194048)!==0)){var a=l.lanes;a&=t.pendingLanes,e|=a,l.lanes=e,wf(t,e)}}function Fi(t,l){var e=t.updateQueue,a=t.alternate;if(a!==null&&(a=a.updateQueue,e===a)){var u=null,n=null;if(e=e.firstBaseUpdate,e!==null){do{var i={lane:e.lane,tag:e.tag,payload:e.payload,callback:null,next:null};n===null?u=n=i:n=n.next=i,e=e.next}while(e!==null);n===null?u=n=l:n=n.next=l}else u=n=l;e={baseState:a.baseState,firstBaseUpdate:u,lastBaseUpdate:n,shared:a.shared,callbacks:a.callbacks},t.updateQueue=e;return}t=e.lastBaseUpdate,t===null?e.firstBaseUpdate=l:t.next=l,e.lastBaseUpdate=l}var Ii=!1;function nu(){if(Ii){var t=ga;if(t!==null)throw t}}function iu(t,l,e,a){Ii=!1;var u=t.updateQueue;oe=!1;var n=u.firstBaseUpdate,i=u.lastBaseUpdate,c=u.shared.pending;if(c!==null){u.shared.pending=null;var s=c,h=s.next;s.next=null,i===null?n=h:i.next=h,i=s;var z=t.alternate;z!==null&&(z=z.updateQueue,c=z.lastBaseUpdate,c!==i&&(c===null?z.firstBaseUpdate=h:c.next=h,z.lastBaseUpdate=s))}if(n!==null){var A=u.baseState;i=0,z=h=s=null,c=n;do{var g=c.lane&-536870913,b=g!==c.lane;if(b?(et&g)===g:(a&g)===g){g!==0&&g===ha&&(Ii=!0),z!==null&&(z=z.next={lane:0,tag:c.tag,payload:c.payload,callback:null,next:null});t:{var B=t,w=c;g=l;var pt=e;switch(w.tag){case 1:if(B=w.payload,typeof B=="function"){A=B.call(pt,A,g);break t}A=B;break t;case 3:B.flags=B.flags&-65537|128;case 0:if(B=w.payload,g=typeof B=="function"?B.call(pt,A,g):B,g==null)break t;A=H({},A,g);break t;case 2:oe=!0}}g=c.callback,g!==null&&(t.flags|=64,b&&(t.flags|=8192),b=u.callbacks,b===null?u.callbacks=[g]:b.push(g))}else b={lane:g,tag:c.tag,payload:c.payload,callback:c.callback,next:null},z===null?(h=z=b,s=A):z=z.next=b,i|=g;if(c=c.next,c===null){if(c=u.shared.pending,c===null)break;b=c,c=b.next,b.next=null,u.lastBaseUpdate=b,u.shared.pending=null}}while(!0);z===null&&(s=A),u.baseState=s,u.firstBaseUpdate=h,u.lastBaseUpdate=z,n===null&&(u.shared.lanes=0),be|=i,t.lanes=i,t.memoizedState=A}}function Pr(t,l){if(typeof t!="function")throw Error(f(191,t));t.call(l)}function ts(t,l){var e=t.callbacks;if(e!==null)for(t.callbacks=null,t=0;tn?n:8;var i=x.T,c={};x.T=c,vc(t,!1,l,e);try{var s=u(),h=x.S;if(h!==null&&h(c,s),s!==null&&typeof s=="object"&&typeof s.then=="function"){var z=xy(s,a);ru(t,l,z,bl(t))}else ru(t,l,a,bl(t))}catch(A){ru(t,l,{then:function(){},status:"rejected",reason:A},bl())}finally{C.p=n,i!==null&&c.types!==null&&(i.types=c.types),x.T=i}}function _y(){}function hc(t,l,e,a){if(t.tag!==5)throw Error(f(476));var u=Cs(t).queue;Us(t,u,l,Z,e===null?_y:function(){return js(t),e(a)})}function Cs(t){var l=t.memoizedState;if(l!==null)return l;l={memoizedState:Z,baseState:Z,baseQueue:null,queue:{pending:null,lanes:0,dispatch:null,lastRenderedReducer:Jl,lastRenderedState:Z},next:null};var e={};return l.next={memoizedState:e,baseState:e,baseQueue:null,queue:{pending:null,lanes:0,dispatch:null,lastRenderedReducer:Jl,lastRenderedState:e},next:null},t.memoizedState=l,t=t.alternate,t!==null&&(t.memoizedState=l),l}function js(t){var l=Cs(t);l.next===null&&(l=t.alternate.memoizedState),ru(t,l.next.queue,{},bl())}function gc(){return Kt(Mu)}function Rs(){return Rt().memoizedState}function Hs(){return Rt().memoizedState}function Oy(t){for(var l=t.return;l!==null;){switch(l.tag){case 24:case 3:var e=bl();t=de(e);var a=ye(l,t,e);a!==null&&(fl(a,l,e),uu(a,l,e)),l={cache:Vi()},t.payload=l;return}l=l.return}}function Ny(t,l,e){var a=bl();e={lane:a,revertLane:0,gesture:null,action:e,hasEagerState:!1,eagerState:null,next:null},Sn(t)?qs(l,e):(e=Ri(t,l,e,a),e!==null&&(fl(e,t,a),Ys(e,l,a)))}function Bs(t,l,e){var a=bl();ru(t,l,e,a)}function ru(t,l,e,a){var u={lane:a,revertLane:0,gesture:null,action:e,hasEagerState:!1,eagerState:null,next:null};if(Sn(t))qs(l,u);else{var n=t.alternate;if(t.lanes===0&&(n===null||n.lanes===0)&&(n=l.lastRenderedReducer,n!==null))try{var i=l.lastRenderedState,c=n(i,e);if(u.hasEagerState=!0,u.eagerState=c,dl(c,i))return Pu(t,l,u,0),St===null&&Iu(),!1}catch{}if(e=Ri(t,l,u,a),e!==null)return fl(e,t,a),Ys(e,l,a),!0}return!1}function vc(t,l,e,a){if(a={lane:2,revertLane:Wc(),gesture:null,action:a,hasEagerState:!1,eagerState:null,next:null},Sn(t)){if(l)throw Error(f(479))}else l=Ri(t,e,a,2),l!==null&&fl(l,t,2)}function Sn(t){var l=t.alternate;return t===W||l!==null&&l===W}function qs(t,l){Sa=yn=!0;var e=t.pending;e===null?l.next=l:(l.next=e.next,e.next=l),t.pending=l}function Ys(t,l,e){if((e&4194048)!==0){var a=l.lanes;a&=t.pendingLanes,e|=a,l.lanes=e,wf(t,e)}}var su={readContext:Kt,use:gn,useCallback:Dt,useContext:Dt,useEffect:Dt,useImperativeHandle:Dt,useLayoutEffect:Dt,useInsertionEffect:Dt,useMemo:Dt,useReducer:Dt,useRef:Dt,useState:Dt,useDebugValue:Dt,useDeferredValue:Dt,useTransition:Dt,useSyncExternalStore:Dt,useId:Dt,useHostTransitionStatus:Dt,useFormState:Dt,useActionState:Dt,useOptimistic:Dt,useMemoCache:Dt,useCacheRefresh:Dt};su.useEffectEvent=Dt;var Gs={readContext:Kt,use:gn,useCallback:function(t,l){return Ft().memoizedState=[t,l===void 0?null:l],t},useContext:Kt,useEffect:zs,useImperativeHandle:function(t,l,e){e=e!=null?e.concat([t]):null,bn(4194308,4,Ms.bind(null,l,t),e)},useLayoutEffect:function(t,l){return bn(4194308,4,t,l)},useInsertionEffect:function(t,l){bn(4,2,t,l)},useMemo:function(t,l){var e=Ft();l=l===void 0?null:l;var a=t();if(Ke){ue(!0);try{t()}finally{ue(!1)}}return e.memoizedState=[a,l],a},useReducer:function(t,l,e){var a=Ft();if(e!==void 0){var u=e(l);if(Ke){ue(!0);try{e(l)}finally{ue(!1)}}}else u=l;return a.memoizedState=a.baseState=u,t={pending:null,lanes:0,dispatch:null,lastRenderedReducer:t,lastRenderedState:u},a.queue=t,t=t.dispatch=Ny.bind(null,W,t),[a.memoizedState,t]},useRef:function(t){var l=Ft();return t={current:t},l.memoizedState=t},useState:function(t){t=sc(t);var l=t.queue,e=Bs.bind(null,W,l);return l.dispatch=e,[t.memoizedState,e]},useDebugValue:yc,useDeferredValue:function(t,l){var e=Ft();return mc(e,t,l)},useTransition:function(){var t=sc(!1);return t=Us.bind(null,W,t.queue,!0,!1),Ft().memoizedState=t,[!1,t]},useSyncExternalStore:function(t,l,e){var a=W,u=Ft();if(ut){if(e===void 0)throw Error(f(407));e=e()}else{if(e=l(),St===null)throw Error(f(349));(et&127)!==0||is(a,l,e)}u.memoizedState=e;var n={value:e,getSnapshot:l};return u.queue=n,zs(fs.bind(null,a,n,t),[t]),a.flags|=2048,za(9,{destroy:void 0},cs.bind(null,a,n,e,l),null),e},useId:function(){var t=Ft(),l=St.identifierPrefix;if(ut){var e=Bl,a=Hl;e=(a&~(1<<32-ol(a)-1)).toString(32)+e,l="_"+l+"R_"+e,e=mn++,0<\/script>",n=n.removeChild(n.firstChild);break;case"select":n=typeof a.is=="string"?i.createElement("select",{is:a.is}):i.createElement("select"),a.multiple?n.multiple=!0:a.size&&(n.size=a.size);break;default:n=typeof a.is=="string"?i.createElement(u,{is:a.is}):i.createElement(u)}}n[Lt]=l,n[el]=a;t:for(i=l.child;i!==null;){if(i.tag===5||i.tag===6)n.appendChild(i.stateNode);else if(i.tag!==4&&i.tag!==27&&i.child!==null){i.child.return=i,i=i.child;continue}if(i===l)break t;for(;i.sibling===null;){if(i.return===null||i.return===l)break t;i=i.return}i.sibling.return=i.return,i=i.sibling}l.stateNode=n;t:switch(kt(n,u,a),u){case"button":case"input":case"select":case"textarea":a=!!a.autoFocus;break t;case"img":a=!0;break t;default:a=!1}a&&Wl(l)}}return At(l),Uc(l,l.type,t===null?null:t.memoizedProps,l.pendingProps,e),null;case 6:if(t&&l.stateNode!=null)t.memoizedProps!==a&&Wl(l);else{if(typeof a!="string"&&l.stateNode===null)throw Error(f(166));if(t=P.current,ya(l)){if(t=l.stateNode,e=l.memoizedProps,a=null,u=Vt,u!==null)switch(u.tag){case 27:case 5:a=u.memoizedProps}t[Lt]=l,t=!!(t.nodeValue===e||a!==null&&a.suppressHydrationWarning===!0||nd(t.nodeValue,e)),t||re(l,!0)}else t=Xn(t).createTextNode(a),t[Lt]=l,l.stateNode=t}return At(l),null;case 31:if(e=l.memoizedState,t===null||t.memoizedState!==null){if(a=ya(l),e!==null){if(t===null){if(!a)throw Error(f(318));if(t=l.memoizedState,t=t!==null?t.dehydrated:null,!t)throw Error(f(557));t[Lt]=l}else Ge(),(l.flags&128)===0&&(l.memoizedState=null),l.flags|=4;At(l),t=!1}else e=Qi(),t!==null&&t.memoizedState!==null&&(t.memoizedState.hydrationErrors=e),t=!0;if(!t)return l.flags&256?(hl(l),l):(hl(l),null);if((l.flags&128)!==0)throw Error(f(558))}return At(l),null;case 13:if(a=l.memoizedState,t===null||t.memoizedState!==null&&t.memoizedState.dehydrated!==null){if(u=ya(l),a!==null&&a.dehydrated!==null){if(t===null){if(!u)throw Error(f(318));if(u=l.memoizedState,u=u!==null?u.dehydrated:null,!u)throw Error(f(317));u[Lt]=l}else Ge(),(l.flags&128)===0&&(l.memoizedState=null),l.flags|=4;At(l),u=!1}else u=Qi(),t!==null&&t.memoizedState!==null&&(t.memoizedState.hydrationErrors=u),u=!0;if(!u)return l.flags&256?(hl(l),l):(hl(l),null)}return hl(l),(l.flags&128)!==0?(l.lanes=e,l):(e=a!==null,t=t!==null&&t.memoizedState!==null,e&&(a=l.child,u=null,a.alternate!==null&&a.alternate.memoizedState!==null&&a.alternate.memoizedState.cachePool!==null&&(u=a.alternate.memoizedState.cachePool.pool),n=null,a.memoizedState!==null&&a.memoizedState.cachePool!==null&&(n=a.memoizedState.cachePool.pool),n!==u&&(a.flags|=2048)),e!==t&&e&&(l.child.flags|=8192),En(l,l.updateQueue),At(l),null);case 4:return Ct(),t===null&&Pc(l.stateNode.containerInfo),At(l),null;case 10:return Vl(l.type),At(l),null;case 19:if(M(jt),a=l.memoizedState,a===null)return At(l),null;if(u=(l.flags&128)!==0,n=a.rendering,n===null)if(u)du(a,!1);else{if(Ut!==0||t!==null&&(t.flags&128)!==0)for(t=l.child;t!==null;){if(n=dn(t),n!==null){for(l.flags|=128,du(a,!1),t=n.updateQueue,l.updateQueue=t,En(l,t),l.subtreeFlags=0,t=e,e=l.child;e!==null;)Br(e,t),e=e.sibling;return j(jt,jt.current&1|2),ut&&wl(l,a.treeForkCount),l.child}t=t.sibling}a.tail!==null&&rl()>Dn&&(l.flags|=128,u=!0,du(a,!1),l.lanes=4194304)}else{if(!u)if(t=dn(n),t!==null){if(l.flags|=128,u=!0,t=t.updateQueue,l.updateQueue=t,En(l,t),du(a,!0),a.tail===null&&a.tailMode==="hidden"&&!n.alternate&&!ut)return At(l),null}else 2*rl()-a.renderingStartTime>Dn&&e!==536870912&&(l.flags|=128,u=!0,du(a,!1),l.lanes=4194304);a.isBackwards?(n.sibling=l.child,l.child=n):(t=a.last,t!==null?t.sibling=n:l.child=n,a.last=n)}return a.tail!==null?(t=a.tail,a.rendering=t,a.tail=t.sibling,a.renderingStartTime=rl(),t.sibling=null,e=jt.current,j(jt,u?e&1|2:e&1),ut&&wl(l,a.treeForkCount),t):(At(l),null);case 22:case 23:return hl(l),tc(),a=l.memoizedState!==null,t!==null?t.memoizedState!==null!==a&&(l.flags|=8192):a&&(l.flags|=8192),a?(e&536870912)!==0&&(l.flags&128)===0&&(At(l),l.subtreeFlags&6&&(l.flags|=8192)):At(l),e=l.updateQueue,e!==null&&En(l,e.retryQueue),e=null,t!==null&&t.memoizedState!==null&&t.memoizedState.cachePool!==null&&(e=t.memoizedState.cachePool.pool),a=null,l.memoizedState!==null&&l.memoizedState.cachePool!==null&&(a=l.memoizedState.cachePool.pool),a!==e&&(l.flags|=2048),t!==null&&M(Ze),null;case 24:return e=null,t!==null&&(e=t.memoizedState.cache),l.memoizedState.cache!==e&&(l.flags|=2048),Vl(Ht),At(l),null;case 25:return null;case 30:return null}throw Error(f(156,l.tag))}function Ry(t,l){switch(Gi(l),l.tag){case 1:return t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 3:return Vl(Ht),Ct(),t=l.flags,(t&65536)!==0&&(t&128)===0?(l.flags=t&-65537|128,l):null;case 26:case 27:case 5:return Hu(l),null;case 31:if(l.memoizedState!==null){if(hl(l),l.alternate===null)throw Error(f(340));Ge()}return t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 13:if(hl(l),t=l.memoizedState,t!==null&&t.dehydrated!==null){if(l.alternate===null)throw Error(f(340));Ge()}return t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 19:return M(jt),null;case 4:return Ct(),null;case 10:return Vl(l.type),null;case 22:case 23:return hl(l),tc(),t!==null&&M(Ze),t=l.flags,t&65536?(l.flags=t&-65537|128,l):null;case 24:return Vl(Ht),null;case 25:return null;default:return null}}function ro(t,l){switch(Gi(l),l.tag){case 3:Vl(Ht),Ct();break;case 26:case 27:case 5:Hu(l);break;case 4:Ct();break;case 31:l.memoizedState!==null&&hl(l);break;case 13:hl(l);break;case 19:M(jt);break;case 10:Vl(l.type);break;case 22:case 23:hl(l),tc(),t!==null&&M(Ze);break;case 24:Vl(Ht)}}function yu(t,l){try{var e=l.updateQueue,a=e!==null?e.lastEffect:null;if(a!==null){var u=a.next;e=u;do{if((e.tag&t)===t){a=void 0;var n=e.create,i=e.inst;a=n(),i.destroy=a}e=e.next}while(e!==u)}}catch(c){ht(l,l.return,c)}}function ge(t,l,e){try{var a=l.updateQueue,u=a!==null?a.lastEffect:null;if(u!==null){var n=u.next;a=n;do{if((a.tag&t)===t){var i=a.inst,c=i.destroy;if(c!==void 0){i.destroy=void 0,u=l;var s=e,h=c;try{h()}catch(z){ht(u,s,z)}}}a=a.next}while(a!==n)}}catch(z){ht(l,l.return,z)}}function so(t){var l=t.updateQueue;if(l!==null){var e=t.stateNode;try{ts(l,e)}catch(a){ht(t,t.return,a)}}}function oo(t,l,e){e.props=Je(t.type,t.memoizedProps),e.state=t.memoizedState;try{e.componentWillUnmount()}catch(a){ht(t,l,a)}}function mu(t,l){try{var e=t.ref;if(e!==null){switch(t.tag){case 26:case 27:case 5:var a=t.stateNode;break;case 30:a=t.stateNode;break;default:a=t.stateNode}typeof e=="function"?t.refCleanup=e(a):e.current=a}}catch(u){ht(t,l,u)}}function ql(t,l){var e=t.ref,a=t.refCleanup;if(e!==null)if(typeof a=="function")try{a()}catch(u){ht(t,l,u)}finally{t.refCleanup=null,t=t.alternate,t!=null&&(t.refCleanup=null)}else if(typeof e=="function")try{e(null)}catch(u){ht(t,l,u)}else e.current=null}function yo(t){var l=t.type,e=t.memoizedProps,a=t.stateNode;try{t:switch(l){case"button":case"input":case"select":case"textarea":e.autoFocus&&a.focus();break t;case"img":e.src?a.src=e.src:e.srcSet&&(a.srcset=e.srcSet)}}catch(u){ht(t,t.return,u)}}function Cc(t,l,e){try{var a=t.stateNode;em(a,t.type,e,l),a[el]=l}catch(u){ht(t,t.return,u)}}function mo(t){return t.tag===5||t.tag===3||t.tag===26||t.tag===27&&Te(t.type)||t.tag===4}function jc(t){t:for(;;){for(;t.sibling===null;){if(t.return===null||mo(t.return))return null;t=t.return}for(t.sibling.return=t.return,t=t.sibling;t.tag!==5&&t.tag!==6&&t.tag!==18;){if(t.tag===27&&Te(t.type)||t.flags&2||t.child===null||t.tag===4)continue t;t.child.return=t,t=t.child}if(!(t.flags&2))return t.stateNode}}function Rc(t,l,e){var a=t.tag;if(a===5||a===6)t=t.stateNode,l?(e.nodeType===9?e.body:e.nodeName==="HTML"?e.ownerDocument.body:e).insertBefore(t,l):(l=e.nodeType===9?e.body:e.nodeName==="HTML"?e.ownerDocument.body:e,l.appendChild(t),e=e._reactRootContainer,e!=null||l.onclick!==null||(l.onclick=Xl));else if(a!==4&&(a===27&&Te(t.type)&&(e=t.stateNode,l=null),t=t.child,t!==null))for(Rc(t,l,e),t=t.sibling;t!==null;)Rc(t,l,e),t=t.sibling}function Mn(t,l,e){var a=t.tag;if(a===5||a===6)t=t.stateNode,l?e.insertBefore(t,l):e.appendChild(t);else if(a!==4&&(a===27&&Te(t.type)&&(e=t.stateNode),t=t.child,t!==null))for(Mn(t,l,e),t=t.sibling;t!==null;)Mn(t,l,e),t=t.sibling}function ho(t){var l=t.stateNode,e=t.memoizedProps;try{for(var a=t.type,u=l.attributes;u.length;)l.removeAttributeNode(u[0]);kt(l,a,e),l[Lt]=t,l[el]=e}catch(n){ht(t,t.return,n)}}var $l=!1,Yt=!1,Hc=!1,go=typeof WeakSet=="function"?WeakSet:Set,Zt=null;function Hy(t,l){if(t=t.containerInfo,ef=Jn,t=_r(t),Oi(t)){if("selectionStart"in t)var e={start:t.selectionStart,end:t.selectionEnd};else t:{e=(e=t.ownerDocument)&&e.defaultView||window;var a=e.getSelection&&e.getSelection();if(a&&a.rangeCount!==0){e=a.anchorNode;var u=a.anchorOffset,n=a.focusNode;a=a.focusOffset;try{e.nodeType,n.nodeType}catch{e=null;break t}var i=0,c=-1,s=-1,h=0,z=0,A=t,g=null;l:for(;;){for(var b;A!==e||u!==0&&A.nodeType!==3||(c=i+u),A!==n||a!==0&&A.nodeType!==3||(s=i+a),A.nodeType===3&&(i+=A.nodeValue.length),(b=A.firstChild)!==null;)g=A,A=b;for(;;){if(A===t)break l;if(g===e&&++h===u&&(c=i),g===n&&++z===a&&(s=i),(b=A.nextSibling)!==null)break;A=g,g=A.parentNode}A=b}e=c===-1||s===-1?null:{start:c,end:s}}else e=null}e=e||{start:0,end:0}}else e=null;for(af={focusedElem:t,selectionRange:e},Jn=!1,Zt=l;Zt!==null;)if(l=Zt,t=l.child,(l.subtreeFlags&1028)!==0&&t!==null)t.return=l,Zt=t;else for(;Zt!==null;){switch(l=Zt,n=l.alternate,t=l.flags,l.tag){case 0:if((t&4)!==0&&(t=l.updateQueue,t=t!==null?t.events:null,t!==null))for(e=0;e title"))),kt(n,a,e),n[Lt]=t,Qt(n),a=n;break t;case"link":var i=zd("link","href",u).get(a+(e.href||""));if(i){for(var c=0;cpt&&(i=pt,pt=w,w=i);var y=Er(c,w),d=Er(c,pt);if(y&&d&&(b.rangeCount!==1||b.anchorNode!==y.node||b.anchorOffset!==y.offset||b.focusNode!==d.node||b.focusOffset!==d.offset)){var m=A.createRange();m.setStart(y.node,y.offset),b.removeAllRanges(),w>pt?(b.addRange(m),b.extend(d.node,d.offset)):(m.setEnd(d.node,d.offset),b.addRange(m))}}}}for(A=[],b=c;b=b.parentNode;)b.nodeType===1&&A.push({element:b,left:b.scrollLeft,top:b.scrollTop});for(typeof c.focus=="function"&&c.focus(),c=0;ce?32:e,x.T=null,e=Zc,Zc=null;var n=Se,i=le;if(Gt=0,_a=Se=null,le=0,(st&6)!==0)throw Error(f(331));var c=st;if(st|=4,_o(n.current),Ao(n,n.current,i,e),st=c,Su(0,!1),sl&&typeof sl.onPostCommitFiberRoot=="function")try{sl.onPostCommitFiberRoot(Ya,n)}catch{}return!0}finally{C.p=u,x.T=a,Vo(t,l)}}function Jo(t,l,e){l=Tl(e,l),l=xc(t.stateNode,l,2),t=ye(t,l,2),t!==null&&(Xa(t,2),Yl(t))}function ht(t,l,e){if(t.tag===3)Jo(t,t,e);else for(;l!==null;){if(l.tag===3){Jo(l,t,e);break}else if(l.tag===1){var a=l.stateNode;if(typeof l.type.getDerivedStateFromError=="function"||typeof a.componentDidCatch=="function"&&(pe===null||!pe.has(a))){t=Tl(e,t),e=Js(2),a=ye(l,e,2),a!==null&&(ks(e,a,l,t),Xa(a,2),Yl(a));break}}l=l.return}}function Kc(t,l,e){var a=t.pingCache;if(a===null){a=t.pingCache=new Yy;var u=new Set;a.set(l,u)}else u=a.get(l),u===void 0&&(u=new Set,a.set(l,u));u.has(e)||(Yc=!0,u.add(e),t=wy.bind(null,t,l,e),l.then(t,t))}function wy(t,l,e){var a=t.pingCache;a!==null&&a.delete(l),t.pingedLanes|=t.suspendedLanes&e,t.warmLanes&=~e,St===t&&(et&e)===e&&(Ut===4||Ut===3&&(et&62914560)===et&&300>rl()-Nn?(st&2)===0&&Oa(t,0):Gc|=e,Ma===et&&(Ma=0)),Yl(t)}function ko(t,l){l===0&&(l=Qf()),t=qe(t,l),t!==null&&(Xa(t,l),Yl(t))}function Ly(t){var l=t.memoizedState,e=0;l!==null&&(e=l.retryLane),ko(t,e)}function Vy(t,l){var e=0;switch(t.tag){case 31:case 13:var a=t.stateNode,u=t.memoizedState;u!==null&&(e=u.retryLane);break;case 19:a=t.stateNode;break;case 22:a=t.stateNode._retryCache;break;default:throw Error(f(314))}a!==null&&a.delete(l),ko(t,e)}function Ky(t,l){return ni(t,l)}var Bn=null,Da=null,Jc=!1,qn=!1,kc=!1,ze=0;function Yl(t){t!==Da&&t.next===null&&(Da===null?Bn=Da=t:Da=Da.next=t),qn=!0,Jc||(Jc=!0,ky())}function Su(t,l){if(!kc&&qn){kc=!0;do for(var e=!1,a=Bn;a!==null;){if(t!==0){var u=a.pendingLanes;if(u===0)var n=0;else{var i=a.suspendedLanes,c=a.pingedLanes;n=(1<<31-ol(42|t)+1)-1,n&=u&~(i&~c),n=n&201326741?n&201326741|1:n?n|2:0}n!==0&&(e=!0,Io(a,n))}else n=et,n=Xu(a,a===St?n:0,a.cancelPendingCommit!==null||a.timeoutHandle!==-1),(n&3)===0||Ga(a,n)||(e=!0,Io(a,n));a=a.next}while(e);kc=!1}}function Jy(){Wo()}function Wo(){qn=Jc=!1;var t=0;ze!==0&&um()&&(t=ze);for(var l=rl(),e=null,a=Bn;a!==null;){var u=a.next,n=$o(a,l);n===0?(a.next=null,e===null?Bn=u:e.next=u,u===null&&(Da=e)):(e=a,(t!==0||(n&3)!==0)&&(qn=!0)),a=u}Gt!==0&&Gt!==5||Su(t),ze!==0&&(ze=0)}function $o(t,l){for(var e=t.suspendedLanes,a=t.pingedLanes,u=t.expirationTimes,n=t.pendingLanes&-62914561;0c)break;var z=s.transferSize,A=s.initiatorType;z&&id(A)&&(s=s.responseEnd,i+=z*(s"u"?null:document;function bd(t,l,e){var a=Ua;if(a&&typeof l=="string"&&l){var u=xl(l);u='link[rel="'+t+'"][href="'+u+'"]',typeof e=="string"&&(u+='[crossorigin="'+e+'"]'),vd.has(u)||(vd.add(u),t={rel:t,crossOrigin:e,href:l},a.querySelector(u)===null&&(l=a.createElement("link"),kt(l,"link",t),Qt(l),a.head.appendChild(l)))}}function ym(t){ee.D(t),bd("dns-prefetch",t,null)}function mm(t,l){ee.C(t,l),bd("preconnect",t,l)}function hm(t,l,e){ee.L(t,l,e);var a=Ua;if(a&&t&&l){var u='link[rel="preload"][as="'+xl(l)+'"]';l==="image"&&e&&e.imageSrcSet?(u+='[imagesrcset="'+xl(e.imageSrcSet)+'"]',typeof e.imageSizes=="string"&&(u+='[imagesizes="'+xl(e.imageSizes)+'"]')):u+='[href="'+xl(t)+'"]';var n=u;switch(l){case"style":n=Ca(t);break;case"script":n=ja(t)}Nl.has(n)||(t=H({rel:"preload",href:l==="image"&&e&&e.imageSrcSet?void 0:t,as:l},e),Nl.set(n,t),a.querySelector(u)!==null||l==="style"&&a.querySelector(Au(n))||l==="script"&&a.querySelector(Eu(n))||(l=a.createElement("link"),kt(l,"link",t),Qt(l),a.head.appendChild(l)))}}function gm(t,l){ee.m(t,l);var e=Ua;if(e&&t){var a=l&&typeof l.as=="string"?l.as:"script",u='link[rel="modulepreload"][as="'+xl(a)+'"][href="'+xl(t)+'"]',n=u;switch(a){case"audioworklet":case"paintworklet":case"serviceworker":case"sharedworker":case"worker":case"script":n=ja(t)}if(!Nl.has(n)&&(t=H({rel:"modulepreload",href:t},l),Nl.set(n,t),e.querySelector(u)===null)){switch(a){case"audioworklet":case"paintworklet":case"serviceworker":case"sharedworker":case"worker":case"script":if(e.querySelector(Eu(n)))return}a=e.createElement("link"),kt(a,"link",t),Qt(a),e.head.appendChild(a)}}}function vm(t,l,e){ee.S(t,l,e);var a=Ua;if(a&&t){var u=ta(a).hoistableStyles,n=Ca(t);l=l||"default";var i=u.get(n);if(!i){var c={loading:0,preload:null};if(i=a.querySelector(Au(n)))c.loading=5;else{t=H({rel:"stylesheet",href:t,"data-precedence":l},e),(e=Nl.get(n))&&of(t,e);var s=i=a.createElement("link");Qt(s),kt(s,"link",t),s._p=new Promise(function(h,z){s.onload=h,s.onerror=z}),s.addEventListener("load",function(){c.loading|=1}),s.addEventListener("error",function(){c.loading|=2}),c.loading|=4,Zn(i,l,a)}i={type:"stylesheet",instance:i,count:1,state:c},u.set(n,i)}}}function bm(t,l){ee.X(t,l);var e=Ua;if(e&&t){var a=ta(e).hoistableScripts,u=ja(t),n=a.get(u);n||(n=e.querySelector(Eu(u)),n||(t=H({src:t,async:!0},l),(l=Nl.get(u))&&df(t,l),n=e.createElement("script"),Qt(n),kt(n,"link",t),e.head.appendChild(n)),n={type:"script",instance:n,count:1,state:null},a.set(u,n))}}function pm(t,l){ee.M(t,l);var e=Ua;if(e&&t){var a=ta(e).hoistableScripts,u=ja(t),n=a.get(u);n||(n=e.querySelector(Eu(u)),n||(t=H({src:t,async:!0,type:"module"},l),(l=Nl.get(u))&&df(t,l),n=e.createElement("script"),Qt(n),kt(n,"link",t),e.head.appendChild(n)),n={type:"script",instance:n,count:1,state:null},a.set(u,n))}}function pd(t,l,e,a){var u=(u=P.current)?Qn(u):null;if(!u)throw Error(f(446));switch(t){case"meta":case"title":return null;case"style":return typeof e.precedence=="string"&&typeof e.href=="string"?(l=Ca(e.href),e=ta(u).hoistableStyles,a=e.get(l),a||(a={type:"style",instance:null,count:0,state:null},e.set(l,a)),a):{type:"void",instance:null,count:0,state:null};case"link":if(e.rel==="stylesheet"&&typeof e.href=="string"&&typeof e.precedence=="string"){t=Ca(e.href);var n=ta(u).hoistableStyles,i=n.get(t);if(i||(u=u.ownerDocument||u,i={type:"stylesheet",instance:null,count:0,state:{loading:0,preload:null}},n.set(t,i),(n=u.querySelector(Au(t)))&&!n._p&&(i.instance=n,i.state.loading=5),Nl.has(t)||(e={rel:"preload",as:"style",href:e.href,crossOrigin:e.crossOrigin,integrity:e.integrity,media:e.media,hrefLang:e.hrefLang,referrerPolicy:e.referrerPolicy},Nl.set(t,e),n||Sm(u,t,e,i.state))),l&&a===null)throw Error(f(528,""));return i}if(l&&a!==null)throw Error(f(529,""));return null;case"script":return l=e.async,e=e.src,typeof e=="string"&&l&&typeof l!="function"&&typeof l!="symbol"?(l=ja(e),e=ta(u).hoistableScripts,a=e.get(l),a||(a={type:"script",instance:null,count:0,state:null},e.set(l,a)),a):{type:"void",instance:null,count:0,state:null};default:throw Error(f(444,t))}}function Ca(t){return'href="'+xl(t)+'"'}function Au(t){return'link[rel="stylesheet"]['+t+"]"}function Sd(t){return H({},t,{"data-precedence":t.precedence,precedence:null})}function Sm(t,l,e,a){t.querySelector('link[rel="preload"][as="style"]['+l+"]")?a.loading=1:(l=t.createElement("link"),a.preload=l,l.addEventListener("load",function(){return a.loading|=1}),l.addEventListener("error",function(){return a.loading|=2}),kt(l,"link",e),Qt(l),t.head.appendChild(l))}function ja(t){return'[src="'+xl(t)+'"]'}function Eu(t){return"script[async]"+t}function xd(t,l,e){if(l.count++,l.instance===null)switch(l.type){case"style":var a=t.querySelector('style[data-href~="'+xl(e.href)+'"]');if(a)return l.instance=a,Qt(a),a;var u=H({},e,{"data-href":e.href,"data-precedence":e.precedence,href:null,precedence:null});return a=(t.ownerDocument||t).createElement("style"),Qt(a),kt(a,"style",u),Zn(a,e.precedence,t),l.instance=a;case"stylesheet":u=Ca(e.href);var n=t.querySelector(Au(u));if(n)return l.state.loading|=4,l.instance=n,Qt(n),n;a=Sd(e),(u=Nl.get(u))&&of(a,u),n=(t.ownerDocument||t).createElement("link"),Qt(n);var i=n;return i._p=new Promise(function(c,s){i.onload=c,i.onerror=s}),kt(n,"link",a),l.state.loading|=4,Zn(n,e.precedence,t),l.instance=n;case"script":return n=ja(e.src),(u=t.querySelector(Eu(n)))?(l.instance=u,Qt(u),u):(a=e,(u=Nl.get(n))&&(a=H({},e),df(a,u)),t=t.ownerDocument||t,u=t.createElement("script"),Qt(u),kt(u,"link",a),t.head.appendChild(u),l.instance=u);case"void":return null;default:throw Error(f(443,l.type))}else l.type==="stylesheet"&&(l.state.loading&4)===0&&(a=l.instance,l.state.loading|=4,Zn(a,e.precedence,t));return l.instance}function Zn(t,l,e){for(var a=e.querySelectorAll('link[rel="stylesheet"][data-precedence],style[data-precedence]'),u=a.length?a[a.length-1]:null,n=u,i=0;i title"):null)}function xm(t,l,e){if(e===1||l.itemProp!=null)return!1;switch(t){case"meta":case"title":return!0;case"style":if(typeof l.precedence!="string"||typeof l.href!="string"||l.href==="")break;return!0;case"link":if(typeof l.rel!="string"||typeof l.href!="string"||l.href===""||l.onLoad||l.onError)break;return l.rel==="stylesheet"?(t=l.disabled,typeof l.precedence=="string"&&t==null):!0;case"script":if(l.async&&typeof l.async!="function"&&typeof l.async!="symbol"&&!l.onLoad&&!l.onError&&l.src&&typeof l.src=="string")return!0}return!1}function Ad(t){return!(t.type==="stylesheet"&&(t.state.loading&3)===0)}function zm(t,l,e,a){if(e.type==="stylesheet"&&(typeof a.media!="string"||matchMedia(a.media).matches!==!1)&&(e.state.loading&4)===0){if(e.instance===null){var u=Ca(a.href),n=l.querySelector(Au(u));if(n){l=n._p,l!==null&&typeof l=="object"&&typeof l.then=="function"&&(t.count++,t=Ln.bind(t),l.then(t,t)),e.state.loading|=4,e.instance=n,Qt(n);return}n=l.ownerDocument||l,a=Sd(a),(u=Nl.get(u))&&of(a,u),n=n.createElement("link"),Qt(n);var i=n;i._p=new Promise(function(c,s){i.onload=c,i.onerror=s}),kt(n,"link",a),e.instance=n}t.stylesheets===null&&(t.stylesheets=new Map),t.stylesheets.set(e,l),(l=e.state.preload)&&(e.state.loading&3)===0&&(t.count++,e=Ln.bind(t),l.addEventListener("load",e),l.addEventListener("error",e))}}var yf=0;function Tm(t,l){return t.stylesheets&&t.count===0&&Kn(t,t.stylesheets),0yf?50:800)+l);return t.unsuspend=e,function(){t.unsuspend=null,clearTimeout(a),clearTimeout(u)}}:null}function Ln(){if(this.count--,this.count===0&&(this.imgCount===0||!this.waitingForImages)){if(this.stylesheets)Kn(this,this.stylesheets);else if(this.unsuspend){var t=this.unsuspend;this.unsuspend=null,t()}}}var Vn=null;function Kn(t,l){t.stylesheets=null,t.unsuspend!==null&&(t.count++,Vn=new Map,l.forEach(Am,t),Vn=null,Ln.call(t))}function Am(t,l){if(!(l.state.loading&4)){var e=Vn.get(t);if(e)var a=e.get(null);else{e=new Map,Vn.set(t,e);for(var u=t.querySelectorAll("link[data-precedence],style[data-precedence]"),n=0;n"u"||typeof __REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE!="function"))try{__REACT_DEVTOOLS_GLOBAL_HOOK__.checkDCE(r)}catch(v){console.error(v)}}return r(),zf.exports=Xm(),zf.exports}var Zm=Qm();const wm=r=>r.replace(/([a-z0-9])([A-Z])/g,"$1-$2").toLowerCase(),Pd=(...r)=>r.filter((v,S,f)=>!!v&&v.trim()!==""&&f.indexOf(v)===S).join(" ").trim();var Lm={xmlns:"http://www.w3.org/2000/svg",width:24,height:24,viewBox:"0 0 24 24",fill:"none",stroke:"currentColor",strokeWidth:2,strokeLinecap:"round",strokeLinejoin:"round"};const Vm=xt.forwardRef(({color:r="currentColor",size:v=24,strokeWidth:S=2,absoluteStrokeWidth:f,className:_="",children:O,iconNode:D,...U},N)=>xt.createElement("svg",{ref:N,...Lm,width:v,height:v,stroke:r,strokeWidth:f?Number(S)*24/Number(v):S,className:Pd("lucide",_),...U},[...D.map(([p,R])=>xt.createElement(p,R)),...Array.isArray(O)?O:[O]]));const Dl=(r,v)=>{const S=xt.forwardRef(({className:f,..._},O)=>xt.createElement(Vm,{ref:O,iconNode:v,className:Pd(`lucide-${wm(r)}`,f),..._}));return S.displayName=`${r}`,S};const Km=Dl("Binary",[["rect",{x:"14",y:"14",width:"4",height:"6",rx:"2",key:"p02svl"}],["rect",{x:"6",y:"4",width:"4",height:"6",rx:"2",key:"xm4xkj"}],["path",{d:"M6 20h4",key:"1i6q5t"}],["path",{d:"M14 10h4",key:"ru81e7"}],["path",{d:"M6 14h2v6",key:"16z9wg"}],["path",{d:"M14 4h2v6",key:"1idq9u"}]]);const Jm=Dl("BookText",[["path",{d:"M4 19.5v-15A2.5 2.5 0 0 1 6.5 2H19a1 1 0 0 1 1 1v18a1 1 0 0 1-1 1H6.5a1 1 0 0 1 0-5H20",key:"k3hazp"}],["path",{d:"M8 11h8",key:"vwpz6n"}],["path",{d:"M8 7h6",key:"1f0q6e"}]]);const km=Dl("EyeOff",[["path",{d:"M10.733 5.076a10.744 10.744 0 0 1 11.205 6.575 1 1 0 0 1 0 .696 10.747 10.747 0 0 1-1.444 2.49",key:"ct8e1f"}],["path",{d:"M14.084 14.158a3 3 0 0 1-4.242-4.242",key:"151rxh"}],["path",{d:"M17.479 17.499a10.75 10.75 0 0 1-15.417-5.151 1 1 0 0 1 0-.696 10.75 10.75 0 0 1 4.446-5.143",key:"13bj9a"}],["path",{d:"m2 2 20 20",key:"1ooewy"}]]);const Wm=Dl("Eye",[["path",{d:"M2.062 12.348a1 1 0 0 1 0-.696 10.75 10.75 0 0 1 19.876 0 1 1 0 0 1 0 .696 10.75 10.75 0 0 1-19.876 0",key:"1nclc0"}],["circle",{cx:"12",cy:"12",r:"3",key:"1v7zrd"}]]);const $m=Dl("Globe",[["circle",{cx:"12",cy:"12",r:"10",key:"1mglay"}],["path",{d:"M12 2a14.5 14.5 0 0 0 0 20 14.5 14.5 0 0 0 0-20",key:"13o1zl"}],["path",{d:"M2 12h20",key:"9i4pu4"}]]);const kd=Dl("LoaderCircle",[["path",{d:"M21 12a9 9 0 1 1-6.219-8.56",key:"13zald"}]]);const Fm=Dl("Lock",[["rect",{width:"18",height:"11",x:"3",y:"11",rx:"2",ry:"2",key:"1w4ew1"}],["path",{d:"M7 11V7a5 5 0 0 1 10 0v4",key:"fwvmzm"}]]);const Im=Dl("LogIn",[["path",{d:"M15 3h4a2 2 0 0 1 2 2v14a2 2 0 0 1-2 2h-4",key:"u53s6r"}],["polyline",{points:"10 17 15 12 10 7",key:"1ail0h"}],["line",{x1:"15",x2:"3",y1:"12",y2:"12",key:"v6grx8"}]]);const Pm=Dl("RotateCw",[["path",{d:"M21 12a9 9 0 1 1-9-9c2.52 0 4.93 1 6.74 2.74L21 8",key:"1p45f6"}],["path",{d:"M21 3v5h-5",key:"1q7to0"}]]);const th=Dl("User",[["path",{d:"M19 21v-2a4 4 0 0 0-4-4H9a4 4 0 0 0-4 4v2",key:"975kel"}],["circle",{cx:"12",cy:"7",r:"4",key:"17ys0d"}]]);const lh=Dl("Waypoints",[["circle",{cx:"12",cy:"4.5",r:"2.5",key:"r5ysbb"}],["path",{d:"m10.2 6.3-3.9 3.9",key:"1nzqf6"}],["circle",{cx:"4.5",cy:"12",r:"2.5",key:"jydg6v"}],["path",{d:"M7 12h10",key:"b7w52i"}],["circle",{cx:"19.5",cy:"12",r:"2.5",key:"1piiel"}],["path",{d:"m13.8 17.7 3.9-3.9",key:"1wyg1y"}],["circle",{cx:"12",cy:"19.5",r:"2.5",key:"13o1pw"}]]);const eh=Dl("X",[["path",{d:"M18 6 6 18",key:"1bl5f8"}],["path",{d:"m6 6 12 12",key:"d8bk6v"}]]);function t0(){return globalThis.__DATA__??{}}function l0(r){var v,S,f="";if(typeof r=="string"||typeof r=="number")f+=r;else if(typeof r=="object")if(Array.isArray(r)){var _=r.length;for(v=0;v<_;v++)r[v]&&(S=l0(r[v]))&&(f&&(f+=" "),f+=S)}else for(S in r)r[S]&&(f&&(f+=" "),f+=S);return f}function ah(){for(var r,v,S=0,f="",_=arguments.length;S<_;S++)(r=arguments[S])&&(v=l0(r))&&(f&&(f+=" "),f+=v);return f}const Hf="-",uh=r=>{const v=ih(r),{conflictingClassGroups:S,conflictingClassGroupModifiers:f}=r;return{getClassGroupId:D=>{const U=D.split(Hf);return U[0]===""&&U.length!==1&&U.shift(),e0(U,v)||nh(D)},getConflictingClassGroupIds:(D,U)=>{const N=S[D]||[];return U&&f[D]?[...N,...f[D]]:N}}},e0=(r,v)=>{if(r.length===0)return v.classGroupId;const S=r[0],f=v.nextPart.get(S),_=f?e0(r.slice(1),f):void 0;if(_)return _;if(v.validators.length===0)return;const O=r.join(Hf);return v.validators.find(({validator:D})=>D(O))?.classGroupId},Wd=/^\[(.+)\]$/,nh=r=>{if(Wd.test(r)){const v=Wd.exec(r)[1],S=v?.substring(0,v.indexOf(":"));if(S)return"arbitrary.."+S}},ih=r=>{const{theme:v,prefix:S}=r,f={nextPart:new Map,validators:[]};return fh(Object.entries(r.classGroups),S).forEach(([O,D])=>{Df(D,f,O,v)}),f},Df=(r,v,S,f)=>{r.forEach(_=>{if(typeof _=="string"){const O=_===""?v:$d(v,_);O.classGroupId=S;return}if(typeof _=="function"){if(ch(_)){Df(_(f),v,S,f);return}v.validators.push({validator:_,classGroupId:S});return}Object.entries(_).forEach(([O,D])=>{Df(D,$d(v,O),S,f)})})},$d=(r,v)=>{let S=r;return v.split(Hf).forEach(f=>{S.nextPart.has(f)||S.nextPart.set(f,{nextPart:new Map,validators:[]}),S=S.nextPart.get(f)}),S},ch=r=>r.isThemeGetter,fh=(r,v)=>v?r.map(([S,f])=>{const _=f.map(O=>typeof O=="string"?v+O:typeof O=="object"?Object.fromEntries(Object.entries(O).map(([D,U])=>[v+D,U])):O);return[S,_]}):r,rh=r=>{if(r<1)return{get:()=>{},set:()=>{}};let v=0,S=new Map,f=new Map;const _=(O,D)=>{S.set(O,D),v++,v>r&&(v=0,f=S,S=new Map)};return{get(O){let D=S.get(O);if(D!==void 0)return D;if((D=f.get(O))!==void 0)return _(O,D),D},set(O,D){S.has(O)?S.set(O,D):_(O,D)}}},a0="!",sh=r=>{const{separator:v,experimentalParseClassName:S}=r,f=v.length===1,_=v[0],O=v.length,D=U=>{const N=[];let p=0,R=0,H;for(let Q=0;QR?H-R:void 0;return{modifiers:N,hasImportantModifier:ot,baseClassName:ct,maybePostfixModifierPosition:G}};return S?U=>S({className:U,parseClassName:D}):D},oh=r=>{if(r.length<=1)return r;const v=[];let S=[];return r.forEach(f=>{f[0]==="["?(v.push(...S.sort(),f),S=[]):S.push(f)}),v.push(...S.sort()),v},dh=r=>({cache:rh(r.cacheSize),parseClassName:sh(r),...uh(r)}),yh=/\s+/,mh=(r,v)=>{const{parseClassName:S,getClassGroupId:f,getConflictingClassGroupIds:_}=v,O=[],D=r.trim().split(yh);let U="";for(let N=D.length-1;N>=0;N-=1){const p=D[N],{modifiers:R,hasImportantModifier:H,baseClassName:L,maybePostfixModifierPosition:ot}=S(p);let ct=!!ot,G=f(ct?L.substring(0,ot):L);if(!G){if(!ct){U=p+(U.length>0?" "+U:U);continue}if(G=f(L),!G){U=p+(U.length>0?" "+U:U);continue}ct=!1}const Q=oh(R).join(":"),V=H?Q+a0:Q,gt=V+G;if(O.includes(gt))continue;O.push(gt);const zt=_(G,ct);for(let _t=0;_t0?" "+U:U)}return U};function hh(){let r=0,v,S,f="";for(;r{if(typeof r=="string")return r;let v,S="";for(let f=0;fH(R),r());return S=dh(p),f=S.cache.get,_=S.cache.set,O=U,U(N)}function U(N){const p=f(N);if(p)return p;const R=mh(N,S);return _(N,R),R}return function(){return O(hh.apply(null,arguments))}}const Et=r=>{const v=S=>S[r]||[];return v.isThemeGetter=!0,v},n0=/^\[(?:([a-z-]+):)?(.+)\]$/i,vh=/^\d+\/\d+$/,bh=new Set(["px","full","screen"]),ph=/^(\d+(\.\d+)?)?(xs|sm|md|lg|xl)$/,Sh=/\d+(%|px|r?em|[sdl]?v([hwib]|min|max)|pt|pc|in|cm|mm|cap|ch|ex|r?lh|cq(w|h|i|b|min|max))|\b(calc|min|max|clamp)\(.+\)|^0$/,xh=/^(rgba?|hsla?|hwb|(ok)?(lab|lch)|color-mix)\(.+\)$/,zh=/^(inset_)?-?((\d+)?\.?(\d+)[a-z]+|0)_-?((\d+)?\.?(\d+)[a-z]+|0)/,Th=/^(url|image|image-set|cross-fade|element|(repeating-)?(linear|radial|conic)-gradient)\(.+\)$/,ae=r=>Ha(r)||bh.has(r)||vh.test(r),Ne=r=>Ba(r,"length",Uh),Ha=r=>!!r&&!Number.isNaN(Number(r)),Mf=r=>Ba(r,"number",Ha),Cu=r=>!!r&&Number.isInteger(Number(r)),Ah=r=>r.endsWith("%")&&Ha(r.slice(0,-1)),F=r=>n0.test(r),De=r=>ph.test(r),Eh=new Set(["length","size","percentage"]),Mh=r=>Ba(r,Eh,i0),_h=r=>Ba(r,"position",i0),Oh=new Set(["image","url"]),Nh=r=>Ba(r,Oh,jh),Dh=r=>Ba(r,"",Ch),ju=()=>!0,Ba=(r,v,S)=>{const f=n0.exec(r);return f?f[1]?typeof v=="string"?f[1]===v:v.has(f[1]):S(f[2]):!1},Uh=r=>Sh.test(r)&&!xh.test(r),i0=()=>!1,Ch=r=>zh.test(r),jh=r=>Th.test(r),Rh=()=>{const r=Et("colors"),v=Et("spacing"),S=Et("blur"),f=Et("brightness"),_=Et("borderColor"),O=Et("borderRadius"),D=Et("borderSpacing"),U=Et("borderWidth"),N=Et("contrast"),p=Et("grayscale"),R=Et("hueRotate"),H=Et("invert"),L=Et("gap"),ot=Et("gradientColorStops"),ct=Et("gradientColorStopPositions"),G=Et("inset"),Q=Et("margin"),V=Et("opacity"),gt=Et("padding"),zt=Et("saturate"),_t=Et("scale"),nt=Et("sepia"),Ot=Et("skew"),J=Et("space"),Nt=Et("translate"),Xt=()=>["auto","contain","none"],pl=()=>["auto","hidden","clip","visible","scroll"],Pt=()=>["auto",F,v],I=()=>[F,v],Rl=()=>["",ae,Ne],tl=()=>["auto",Ha,F],ll=()=>["bottom","center","left","left-bottom","left-top","right","right-bottom","right-top","top"],x=()=>["solid","dashed","dotted","double","none"],C=()=>["normal","multiply","screen","overlay","darken","lighten","color-dodge","color-burn","hard-light","soft-light","difference","exclusion","hue","saturation","color","luminosity"],Z=()=>["start","end","center","between","around","evenly","stretch"],it=()=>["","0",F],dt=()=>["auto","avoid","all","avoid-page","page","left","right","column"],o=()=>[Ha,F];return{cacheSize:500,separator:":",theme:{colors:[ju],spacing:[ae,Ne],blur:["none","",De,F],brightness:o(),borderColor:[r],borderRadius:["none","","full",De,F],borderSpacing:I(),borderWidth:Rl(),contrast:o(),grayscale:it(),hueRotate:o(),invert:it(),gap:I(),gradientColorStops:[r],gradientColorStopPositions:[Ah,Ne],inset:Pt(),margin:Pt(),opacity:o(),padding:I(),saturate:o(),scale:o(),sepia:it(),skew:o(),space:I(),translate:I()},classGroups:{aspect:[{aspect:["auto","square","video",F]}],container:["container"],columns:[{columns:[De]}],"break-after":[{"break-after":dt()}],"break-before":[{"break-before":dt()}],"break-inside":[{"break-inside":["auto","avoid","avoid-page","avoid-column"]}],"box-decoration":[{"box-decoration":["slice","clone"]}],box:[{box:["border","content"]}],display:["block","inline-block","inline","flex","inline-flex","table","inline-table","table-caption","table-cell","table-column","table-column-group","table-footer-group","table-header-group","table-row-group","table-row","flow-root","grid","inline-grid","contents","list-item","hidden"],float:[{float:["right","left","none","start","end"]}],clear:[{clear:["left","right","both","none","start","end"]}],isolation:["isolate","isolation-auto"],"object-fit":[{object:["contain","cover","fill","none","scale-down"]}],"object-position":[{object:[...ll(),F]}],overflow:[{overflow:pl()}],"overflow-x":[{"overflow-x":pl()}],"overflow-y":[{"overflow-y":pl()}],overscroll:[{overscroll:Xt()}],"overscroll-x":[{"overscroll-x":Xt()}],"overscroll-y":[{"overscroll-y":Xt()}],position:["static","fixed","absolute","relative","sticky"],inset:[{inset:[G]}],"inset-x":[{"inset-x":[G]}],"inset-y":[{"inset-y":[G]}],start:[{start:[G]}],end:[{end:[G]}],top:[{top:[G]}],right:[{right:[G]}],bottom:[{bottom:[G]}],left:[{left:[G]}],visibility:["visible","invisible","collapse"],z:[{z:["auto",Cu,F]}],basis:[{basis:Pt()}],"flex-direction":[{flex:["row","row-reverse","col","col-reverse"]}],"flex-wrap":[{flex:["wrap","wrap-reverse","nowrap"]}],flex:[{flex:["1","auto","initial","none",F]}],grow:[{grow:it()}],shrink:[{shrink:it()}],order:[{order:["first","last","none",Cu,F]}],"grid-cols":[{"grid-cols":[ju]}],"col-start-end":[{col:["auto",{span:["full",Cu,F]},F]}],"col-start":[{"col-start":tl()}],"col-end":[{"col-end":tl()}],"grid-rows":[{"grid-rows":[ju]}],"row-start-end":[{row:["auto",{span:[Cu,F]},F]}],"row-start":[{"row-start":tl()}],"row-end":[{"row-end":tl()}],"grid-flow":[{"grid-flow":["row","col","dense","row-dense","col-dense"]}],"auto-cols":[{"auto-cols":["auto","min","max","fr",F]}],"auto-rows":[{"auto-rows":["auto","min","max","fr",F]}],gap:[{gap:[L]}],"gap-x":[{"gap-x":[L]}],"gap-y":[{"gap-y":[L]}],"justify-content":[{justify:["normal",...Z()]}],"justify-items":[{"justify-items":["start","end","center","stretch"]}],"justify-self":[{"justify-self":["auto","start","end","center","stretch"]}],"align-content":[{content:["normal",...Z(),"baseline"]}],"align-items":[{items:["start","end","center","baseline","stretch"]}],"align-self":[{self:["auto","start","end","center","stretch","baseline"]}],"place-content":[{"place-content":[...Z(),"baseline"]}],"place-items":[{"place-items":["start","end","center","baseline","stretch"]}],"place-self":[{"place-self":["auto","start","end","center","stretch"]}],p:[{p:[gt]}],px:[{px:[gt]}],py:[{py:[gt]}],ps:[{ps:[gt]}],pe:[{pe:[gt]}],pt:[{pt:[gt]}],pr:[{pr:[gt]}],pb:[{pb:[gt]}],pl:[{pl:[gt]}],m:[{m:[Q]}],mx:[{mx:[Q]}],my:[{my:[Q]}],ms:[{ms:[Q]}],me:[{me:[Q]}],mt:[{mt:[Q]}],mr:[{mr:[Q]}],mb:[{mb:[Q]}],ml:[{ml:[Q]}],"space-x":[{"space-x":[J]}],"space-x-reverse":["space-x-reverse"],"space-y":[{"space-y":[J]}],"space-y-reverse":["space-y-reverse"],w:[{w:["auto","min","max","fit","svw","lvw","dvw",F,v]}],"min-w":[{"min-w":[F,v,"min","max","fit"]}],"max-w":[{"max-w":[F,v,"none","full","min","max","fit","prose",{screen:[De]},De]}],h:[{h:[F,v,"auto","min","max","fit","svh","lvh","dvh"]}],"min-h":[{"min-h":[F,v,"min","max","fit","svh","lvh","dvh"]}],"max-h":[{"max-h":[F,v,"min","max","fit","svh","lvh","dvh"]}],size:[{size:[F,v,"auto","min","max","fit"]}],"font-size":[{text:["base",De,Ne]}],"font-smoothing":["antialiased","subpixel-antialiased"],"font-style":["italic","not-italic"],"font-weight":[{font:["thin","extralight","light","normal","medium","semibold","bold","extrabold","black",Mf]}],"font-family":[{font:[ju]}],"fvn-normal":["normal-nums"],"fvn-ordinal":["ordinal"],"fvn-slashed-zero":["slashed-zero"],"fvn-figure":["lining-nums","oldstyle-nums"],"fvn-spacing":["proportional-nums","tabular-nums"],"fvn-fraction":["diagonal-fractions","stacked-fractions"],tracking:[{tracking:["tighter","tight","normal","wide","wider","widest",F]}],"line-clamp":[{"line-clamp":["none",Ha,Mf]}],leading:[{leading:["none","tight","snug","normal","relaxed","loose",ae,F]}],"list-image":[{"list-image":["none",F]}],"list-style-type":[{list:["none","disc","decimal",F]}],"list-style-position":[{list:["inside","outside"]}],"placeholder-color":[{placeholder:[r]}],"placeholder-opacity":[{"placeholder-opacity":[V]}],"text-alignment":[{text:["left","center","right","justify","start","end"]}],"text-color":[{text:[r]}],"text-opacity":[{"text-opacity":[V]}],"text-decoration":["underline","overline","line-through","no-underline"],"text-decoration-style":[{decoration:[...x(),"wavy"]}],"text-decoration-thickness":[{decoration:["auto","from-font",ae,Ne]}],"underline-offset":[{"underline-offset":["auto",ae,F]}],"text-decoration-color":[{decoration:[r]}],"text-transform":["uppercase","lowercase","capitalize","normal-case"],"text-overflow":["truncate","text-ellipsis","text-clip"],"text-wrap":[{text:["wrap","nowrap","balance","pretty"]}],indent:[{indent:I()}],"vertical-align":[{align:["baseline","top","middle","bottom","text-top","text-bottom","sub","super",F]}],whitespace:[{whitespace:["normal","nowrap","pre","pre-line","pre-wrap","break-spaces"]}],break:[{break:["normal","words","all","keep"]}],hyphens:[{hyphens:["none","manual","auto"]}],content:[{content:["none",F]}],"bg-attachment":[{bg:["fixed","local","scroll"]}],"bg-clip":[{"bg-clip":["border","padding","content","text"]}],"bg-opacity":[{"bg-opacity":[V]}],"bg-origin":[{"bg-origin":["border","padding","content"]}],"bg-position":[{bg:[...ll(),_h]}],"bg-repeat":[{bg:["no-repeat",{repeat:["","x","y","round","space"]}]}],"bg-size":[{bg:["auto","cover","contain",Mh]}],"bg-image":[{bg:["none",{"gradient-to":["t","tr","r","br","b","bl","l","tl"]},Nh]}],"bg-color":[{bg:[r]}],"gradient-from-pos":[{from:[ct]}],"gradient-via-pos":[{via:[ct]}],"gradient-to-pos":[{to:[ct]}],"gradient-from":[{from:[ot]}],"gradient-via":[{via:[ot]}],"gradient-to":[{to:[ot]}],rounded:[{rounded:[O]}],"rounded-s":[{"rounded-s":[O]}],"rounded-e":[{"rounded-e":[O]}],"rounded-t":[{"rounded-t":[O]}],"rounded-r":[{"rounded-r":[O]}],"rounded-b":[{"rounded-b":[O]}],"rounded-l":[{"rounded-l":[O]}],"rounded-ss":[{"rounded-ss":[O]}],"rounded-se":[{"rounded-se":[O]}],"rounded-ee":[{"rounded-ee":[O]}],"rounded-es":[{"rounded-es":[O]}],"rounded-tl":[{"rounded-tl":[O]}],"rounded-tr":[{"rounded-tr":[O]}],"rounded-br":[{"rounded-br":[O]}],"rounded-bl":[{"rounded-bl":[O]}],"border-w":[{border:[U]}],"border-w-x":[{"border-x":[U]}],"border-w-y":[{"border-y":[U]}],"border-w-s":[{"border-s":[U]}],"border-w-e":[{"border-e":[U]}],"border-w-t":[{"border-t":[U]}],"border-w-r":[{"border-r":[U]}],"border-w-b":[{"border-b":[U]}],"border-w-l":[{"border-l":[U]}],"border-opacity":[{"border-opacity":[V]}],"border-style":[{border:[...x(),"hidden"]}],"divide-x":[{"divide-x":[U]}],"divide-x-reverse":["divide-x-reverse"],"divide-y":[{"divide-y":[U]}],"divide-y-reverse":["divide-y-reverse"],"divide-opacity":[{"divide-opacity":[V]}],"divide-style":[{divide:x()}],"border-color":[{border:[_]}],"border-color-x":[{"border-x":[_]}],"border-color-y":[{"border-y":[_]}],"border-color-s":[{"border-s":[_]}],"border-color-e":[{"border-e":[_]}],"border-color-t":[{"border-t":[_]}],"border-color-r":[{"border-r":[_]}],"border-color-b":[{"border-b":[_]}],"border-color-l":[{"border-l":[_]}],"divide-color":[{divide:[_]}],"outline-style":[{outline:["",...x()]}],"outline-offset":[{"outline-offset":[ae,F]}],"outline-w":[{outline:[ae,Ne]}],"outline-color":[{outline:[r]}],"ring-w":[{ring:Rl()}],"ring-w-inset":["ring-inset"],"ring-color":[{ring:[r]}],"ring-opacity":[{"ring-opacity":[V]}],"ring-offset-w":[{"ring-offset":[ae,Ne]}],"ring-offset-color":[{"ring-offset":[r]}],shadow:[{shadow:["","inner","none",De,Dh]}],"shadow-color":[{shadow:[ju]}],opacity:[{opacity:[V]}],"mix-blend":[{"mix-blend":[...C(),"plus-lighter","plus-darker"]}],"bg-blend":[{"bg-blend":C()}],filter:[{filter:["","none"]}],blur:[{blur:[S]}],brightness:[{brightness:[f]}],contrast:[{contrast:[N]}],"drop-shadow":[{"drop-shadow":["","none",De,F]}],grayscale:[{grayscale:[p]}],"hue-rotate":[{"hue-rotate":[R]}],invert:[{invert:[H]}],saturate:[{saturate:[zt]}],sepia:[{sepia:[nt]}],"backdrop-filter":[{"backdrop-filter":["","none"]}],"backdrop-blur":[{"backdrop-blur":[S]}],"backdrop-brightness":[{"backdrop-brightness":[f]}],"backdrop-contrast":[{"backdrop-contrast":[N]}],"backdrop-grayscale":[{"backdrop-grayscale":[p]}],"backdrop-hue-rotate":[{"backdrop-hue-rotate":[R]}],"backdrop-invert":[{"backdrop-invert":[H]}],"backdrop-opacity":[{"backdrop-opacity":[V]}],"backdrop-saturate":[{"backdrop-saturate":[zt]}],"backdrop-sepia":[{"backdrop-sepia":[nt]}],"border-collapse":[{border:["collapse","separate"]}],"border-spacing":[{"border-spacing":[D]}],"border-spacing-x":[{"border-spacing-x":[D]}],"border-spacing-y":[{"border-spacing-y":[D]}],"table-layout":[{table:["auto","fixed"]}],caption:[{caption:["top","bottom"]}],transition:[{transition:["none","all","","colors","opacity","shadow","transform",F]}],duration:[{duration:o()}],ease:[{ease:["linear","in","out","in-out",F]}],delay:[{delay:o()}],animate:[{animate:["none","spin","ping","pulse","bounce",F]}],transform:[{transform:["","gpu","none"]}],scale:[{scale:[_t]}],"scale-x":[{"scale-x":[_t]}],"scale-y":[{"scale-y":[_t]}],rotate:[{rotate:[Cu,F]}],"translate-x":[{"translate-x":[Nt]}],"translate-y":[{"translate-y":[Nt]}],"skew-x":[{"skew-x":[Ot]}],"skew-y":[{"skew-y":[Ot]}],"transform-origin":[{origin:["center","top","top-right","right","bottom-right","bottom","bottom-left","left","top-left",F]}],accent:[{accent:["auto",r]}],appearance:[{appearance:["none","auto"]}],cursor:[{cursor:["auto","default","pointer","wait","text","move","help","not-allowed","none","context-menu","progress","cell","crosshair","vertical-text","alias","copy","no-drop","grab","grabbing","all-scroll","col-resize","row-resize","n-resize","e-resize","s-resize","w-resize","ne-resize","nw-resize","se-resize","sw-resize","ew-resize","ns-resize","nesw-resize","nwse-resize","zoom-in","zoom-out",F]}],"caret-color":[{caret:[r]}],"pointer-events":[{"pointer-events":["none","auto"]}],resize:[{resize:["none","y","x",""]}],"scroll-behavior":[{scroll:["auto","smooth"]}],"scroll-m":[{"scroll-m":I()}],"scroll-mx":[{"scroll-mx":I()}],"scroll-my":[{"scroll-my":I()}],"scroll-ms":[{"scroll-ms":I()}],"scroll-me":[{"scroll-me":I()}],"scroll-mt":[{"scroll-mt":I()}],"scroll-mr":[{"scroll-mr":I()}],"scroll-mb":[{"scroll-mb":I()}],"scroll-ml":[{"scroll-ml":I()}],"scroll-p":[{"scroll-p":I()}],"scroll-px":[{"scroll-px":I()}],"scroll-py":[{"scroll-py":I()}],"scroll-ps":[{"scroll-ps":I()}],"scroll-pe":[{"scroll-pe":I()}],"scroll-pt":[{"scroll-pt":I()}],"scroll-pr":[{"scroll-pr":I()}],"scroll-pb":[{"scroll-pb":I()}],"scroll-pl":[{"scroll-pl":I()}],"snap-align":[{snap:["start","end","center","align-none"]}],"snap-stop":[{snap:["normal","always"]}],"snap-type":[{snap:["none","x","y","both"]}],"snap-strictness":[{snap:["mandatory","proximity"]}],touch:[{touch:["auto","none","manipulation"]}],"touch-x":[{"touch-pan":["x","left","right"]}],"touch-y":[{"touch-pan":["y","up","down"]}],"touch-pz":["touch-pinch-zoom"],select:[{select:["none","text","all","auto"]}],"will-change":[{"will-change":["auto","scroll","contents","transform",F]}],fill:[{fill:[r,"none"]}],"stroke-w":[{stroke:[ae,Ne,Mf]}],stroke:[{stroke:[r,"none"]}],sr:["sr-only","not-sr-only"],"forced-color-adjust":[{"forced-color-adjust":["auto","none"]}]},conflictingClassGroups:{overflow:["overflow-x","overflow-y"],overscroll:["overscroll-x","overscroll-y"],inset:["inset-x","inset-y","start","end","top","right","bottom","left"],"inset-x":["right","left"],"inset-y":["top","bottom"],flex:["basis","grow","shrink"],gap:["gap-x","gap-y"],p:["px","py","ps","pe","pt","pr","pb","pl"],px:["pr","pl"],py:["pt","pb"],m:["mx","my","ms","me","mt","mr","mb","ml"],mx:["mr","ml"],my:["mt","mb"],size:["w","h"],"font-size":["leading"],"fvn-normal":["fvn-ordinal","fvn-slashed-zero","fvn-figure","fvn-spacing","fvn-fraction"],"fvn-ordinal":["fvn-normal"],"fvn-slashed-zero":["fvn-normal"],"fvn-figure":["fvn-normal"],"fvn-spacing":["fvn-normal"],"fvn-fraction":["fvn-normal"],"line-clamp":["display","overflow"],rounded:["rounded-s","rounded-e","rounded-t","rounded-r","rounded-b","rounded-l","rounded-ss","rounded-se","rounded-ee","rounded-es","rounded-tl","rounded-tr","rounded-br","rounded-bl"],"rounded-s":["rounded-ss","rounded-es"],"rounded-e":["rounded-se","rounded-ee"],"rounded-t":["rounded-tl","rounded-tr"],"rounded-r":["rounded-tr","rounded-br"],"rounded-b":["rounded-br","rounded-bl"],"rounded-l":["rounded-tl","rounded-bl"],"border-spacing":["border-spacing-x","border-spacing-y"],"border-w":["border-w-s","border-w-e","border-w-t","border-w-r","border-w-b","border-w-l"],"border-w-x":["border-w-r","border-w-l"],"border-w-y":["border-w-t","border-w-b"],"border-color":["border-color-s","border-color-e","border-color-t","border-color-r","border-color-b","border-color-l"],"border-color-x":["border-color-r","border-color-l"],"border-color-y":["border-color-t","border-color-b"],"scroll-m":["scroll-mx","scroll-my","scroll-ms","scroll-me","scroll-mt","scroll-mr","scroll-mb","scroll-ml"],"scroll-mx":["scroll-mr","scroll-ml"],"scroll-my":["scroll-mt","scroll-mb"],"scroll-p":["scroll-px","scroll-py","scroll-ps","scroll-pe","scroll-pt","scroll-pr","scroll-pb","scroll-pl"],"scroll-px":["scroll-pr","scroll-pl"],"scroll-py":["scroll-pt","scroll-pb"],touch:["touch-x","touch-y","touch-pz"],"touch-x":["touch"],"touch-y":["touch"],"touch-pz":["touch"]},conflictingClassGroupModifiers:{"font-size":["leading"]}}},Hh=gh(Rh);function wt(...r){return Hh(ah(r))}const Bh=["relative cursor-pointer","text-sm focus:z-10 focus:ring-2 font-medium focus:outline-none whitespace-nowrap shadow-sm","inline-flex gap-2 items-center justify-center transition-colors focus:ring-offset-1","disabled:opacity-40 disabled:cursor-not-allowed disabled:text-nb-gray-300 ring-offset-neutral-950/50"],qh={default:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:focus:ring-zinc-800/50 dark:bg-nb-gray dark:text-gray-400 dark:border-gray-700/30 dark:hover:text-white dark:hover:bg-zinc-800/50"],primary:["dark:focus:ring-netbird-600/50 dark:ring-offset-neutral-950/50 enabled:dark:bg-netbird disabled:dark:bg-nb-gray-910 dark:text-gray-100 enabled:dark:hover:text-white enabled:dark:hover:bg-netbird-500/80","enabled:bg-netbird enabled:text-white enabled:focus:ring-netbird-400/50 enabled:hover:bg-netbird-500"],secondary:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-920 dark:text-gray-400 dark:border-gray-700/40 dark:hover:text-white dark:hover:bg-nb-gray-910"],secondaryLighter:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900/70 dark:text-gray-400 dark:border-gray-700/70 dark:hover:text-white dark:hover:bg-nb-gray-800/60"],input:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-neutral-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900 dark:text-gray-400 dark:border-nb-gray-700 dark:hover:bg-nb-gray-900/80"],dropdown:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-neutral-200 text-gray-900","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900/40 dark:text-gray-400 dark:border-nb-gray-900 dark:hover:bg-nb-gray-900/50"],dotted:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900 border-dashed","dark:ring-offset-neutral-950/50 dark:focus:ring-neutral-500/20","dark:bg-nb-gray-900/30 dark:text-gray-400 dark:border-gray-500/40 dark:hover:text-white dark:hover:bg-zinc-800/50"],tertiary:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:focus:ring-zinc-800/50 dark:bg-white dark:text-gray-800 dark:border-gray-700/40 dark:hover:bg-neutral-200 disabled:dark:bg-nb-gray-920 disabled:dark:text-nb-gray-300"],white:["focus:ring-white/50 bg-white text-gray-800 border-white outline-none hover:bg-neutral-200 disabled:dark:bg-nb-gray-920 disabled:dark:text-nb-gray-300","disabled:dark:bg-nb-gray-900 disabled:dark:text-nb-gray-300 disabled:dark:border-nb-gray-900"],outline:["bg-white hover:text-black focus:ring-zinc-200/50 hover:bg-gray-100 border-gray-200 text-gray-900","dark:focus:ring-zinc-800/50 dark:bg-transparent dark:text-netbird dark:border-netbird dark:hover:bg-nb-gray-900/30"],"danger-outline":["enabled:dark:focus:ring-red-800/20 enabled:dark:focus:bg-red-950/40 enabled:hover:dark:bg-red-950/50 enabled:dark:hover:border-red-800/50 dark:bg-transparent dark:text-red-500"],"danger-text":["dark:bg-transparent dark:text-red-500 dark:hover:text-red-600 dark:border-transparent !px-0 !shadow-none !py-0 focus:ring-red-500/30 dark:ring-offset-neutral-950/50"],"default-outline":["dark:ring-offset-nb-gray-950/50 dark:focus:ring-nb-gray-500/20","dark:bg-transparent dark:text-nb-gray-400 dark:border-transparent dark:hover:text-white dark:hover:bg-nb-gray-900/30 dark:hover:border-nb-gray-800/50","data-[state=open]:dark:text-white data-[state=open]:dark:bg-nb-gray-900/30 data-[state=open]:dark:border-nb-gray-800/50"],danger:["dark:focus:ring-red-700/20 dark:focus:bg-red-700 hover:dark:bg-red-700 dark:hover:border-red-800/50 dark:bg-red-600 dark:text-red-100"]},Yh={xs:"text-xs py-2 px-4",xs2:"text-[0.78rem] py-2 px-4",sm:"text-sm py-2.5 px-4",md:"text-sm py-2.5 px-4",lg:"text-base py-2.5 px-4"},Gh={0:"border",1:"border border-transparent",2:"border border-t-0 border-b-0"},Ru=xt.forwardRef(({variant:r="default",rounded:v=!0,border:S=1,size:f="md",stopPropagation:_=!0,className:O,onClick:D,children:U,...N},p)=>E.jsx("button",{type:"button",...N,ref:p,className:wt(Bh,qh[r],Yh[f],Gh[S?1:0],v&&"rounded-md",O),onClick:R=>{_&&R.stopPropagation(),D?.(R)},children:U}));Ru.displayName="Button";const Xh={default:["bg-nb-gray-900 placeholder:text-neutral-400/70 border-nb-gray-700","ring-offset-neutral-950/50 focus-visible:ring-neutral-500/20"],darker:["bg-nb-gray-920 placeholder:text-neutral-400/70 border-nb-gray-800","ring-offset-neutral-950/50 focus-visible:ring-neutral-500/20"],error:["bg-nb-gray-900 placeholder:text-neutral-400/70 border-red-500 text-red-500","ring-offset-red-500/10 focus-visible:ring-red-500/10"]},Qh={default:"bg-nb-gray-900 border-nb-gray-700 text-nb-gray-300",error:"bg-nb-gray-900 border-red-500 text-nb-gray-300 text-red-500"},c0=xt.forwardRef(({className:r,type:v,customSuffix:S,customPrefix:f,icon:_,maxWidthClass:O="",error:D,variant:U="default",prefixClassName:N,showPasswordToggle:p=!1,...R},H)=>{const[L,ot]=xt.useState(!1),ct=v==="password",G=ct&&L?"text":v,V=(ct&&p?E.jsx("button",{type:"button",onClick:()=>ot(!L),className:"hover:text-white transition-all","aria-label":"Toggle password visibility",children:L?E.jsx(km,{size:18}):E.jsx(Wm,{size:18})}):null)||S,gt=D?"error":U;return E.jsxs(E.Fragment,{children:[E.jsxs("div",{className:wt("flex relative h-[42px]",O),children:[f&&E.jsx("div",{className:wt(Qh[D?"error":"default"],"flex h-[42px] w-auto rounded-l-md px-3 py-2 text-sm","border items-center whitespace-nowrap",R.disabled&&"opacity-40",N),children:f}),E.jsx("div",{className:wt("absolute left-0 top-0 h-full flex items-center text-xs text-nb-gray-300 pl-3 leading-[0]",R.disabled&&"opacity-40"),children:_}),E.jsx("input",{type:G,ref:H,...R,className:wt(Xh[gt],"flex h-[42px] w-full rounded-md px-3 py-2 text-sm","file:bg-transparent file:text-sm file:font-medium file:border-0","focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-offset-2","disabled:cursor-not-allowed disabled:opacity-40","border",f&&"!border-l-0 !rounded-l-none",V&&"!pr-16",_&&"!pl-10",r)}),E.jsx("div",{className:wt("absolute right-0 top-0 h-full flex items-center text-xs text-nb-gray-300 pr-4 leading-[0] select-none",R.disabled&&"opacity-30"),children:V})]}),D&&E.jsx("p",{className:"text-xs text-red-500 mt-2",children:D})]})});c0.displayName="Input";const Zh=xt.forwardRef(function({value:v,onChange:S,length:f=6,disabled:_=!1,className:O,autoFocus:D=!1},U){const N=xt.useRef([]);xt.useImperativeHandle(U,()=>({focus:()=>{N.current[0]?.focus()}}));const p=v.split("").concat(new Array(f).fill("")).slice(0,f),R=Array.from({length:f},(G,Q)=>`pin-${Q}`),H=(G,Q)=>{if(!/^\d*$/.test(Q))return;const V=[...p];V[G]=Q.slice(-1);const gt=V.join("").replaceAll(/\s/g,"");S(gt),Q&&G{Q.key==="Backspace"&&!p[G]&&G>0&&N.current[G-1]?.focus(),Q.key==="ArrowLeft"&&G>0&&N.current[G-1]?.focus(),Q.key==="ArrowRight"&&G{G.preventDefault();const Q=G.clipboardData.getData("text").replaceAll(/\D/g,"").slice(0,f);S(Q);const V=Math.min(Q.length,f-1);N.current[V]?.focus()},ct=G=>{G.target.select()};return E.jsx("div",{className:wt("flex gap-2 w-full min-w-0",O),children:p.map((G,Q)=>E.jsx("input",{id:R[Q],ref:V=>{N.current[Q]=V},type:"text",inputMode:"numeric",maxLength:1,value:G,onChange:V=>H(Q,V.target.value),onKeyDown:V=>L(Q,V),onPaste:ot,onFocus:ct,disabled:_,autoFocus:D&&Q===0,className:wt("flex-1 min-w-0 h-[42px] text-center text-sm rounded-md","dark:bg-nb-gray-900 border dark:border-nb-gray-700","dark:placeholder:text-neutral-400/70","focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-offset-2","ring-offset-neutral-200/20 dark:ring-offset-neutral-950/50 dark:focus-visible:ring-neutral-500/20","disabled:cursor-not-allowed disabled:opacity-40")},R[Q]))})}),f0=xt.createContext({value:"",onChange:()=>{}}),r0=()=>xt.useContext(f0);function $e({value:r,defaultValue:v,onChange:S,children:f}){const[_,O]=xt.useState(v??""),D=r??_,U=xt.useCallback(p=>{r===void 0&&O(p),S?.(p)},[r,S]),N=xt.useMemo(()=>({value:D,onChange:U}),[D,U]);return E.jsx(f0.Provider,{value:N,children:E.jsx("div",{children:typeof f=="function"?f({value:D,onChange:U}):f})})}function wh({children:r,className:v}){return E.jsx("div",{role:"tablist",className:wt("bg-nb-gray-930/70 p-1.5 flex justify-center gap-1 border-nb-gray-900",v),children:r})}function Lh({children:r,value:v,disabled:S=!1,className:f,selected:_,onClick:O}){const D=r0(),U=_??v===D.value;let N="";U?N="bg-nb-gray-900 text-white":S||(N="text-nb-gray-400 hover:bg-nb-gray-900/50");const p=()=>{D.onChange(v),O?.()};return E.jsx("button",{role:"tab",type:"button",disabled:S,"aria-selected":U,onClick:p,className:wt("px-4 py-2 text-sm rounded-md w-full transition-all cursor-pointer",S&&"opacity-30 cursor-not-allowed",N,f),children:E.jsx("div",{className:"flex items-center w-full justify-center gap-2",children:r})})}function Vh({children:r,value:v,className:S,visible:f}){const _=r0();return f??v===_.value?E.jsx("div",{role:"tabpanel",className:wt("bg-nb-gray-930/70 px-4 pt-4 pb-5 rounded-b-md border border-t-0 border-nb-gray-900",S),children:r}):null}$e.List=wh;$e.Trigger=Lh;$e.Content=Vh;const Kh="/__netbird__/assets/netbird-full.svg",Jh="data:image/svg+xml,%3csvg%20width='31'%20height='23'%20viewBox='0%200%2031%2023'%20fill='none'%20xmlns='http://www.w3.org/2000/svg'%3e%3cpath%20d='M21.4631%200.523438C17.8173%200.857913%2016.0028%202.95675%2015.3171%204.01871L4.66406%2022.4734H17.5163L30.1929%200.523438H21.4631Z'%20fill='%23F68330'/%3e%3cpath%20d='M17.5265%2022.4737L0%203.88525C0%203.88525%2019.8177%20-1.44128%2021.7493%2015.1738L17.5265%2022.4737Z'%20fill='%23F68330'/%3e%3cpath%20d='M14.9236%204.70563L9.54688%2014.0208L17.5158%2022.4747L21.7385%2015.158C21.0696%209.44682%2018.2851%206.32784%2014.9236%204.69727'%20fill='%23F05252'/%3e%3c/svg%3e",ti={small:{desktop:14,mobile:20},default:{desktop:22,mobile:30},large:{desktop:24,mobile:40}},kh=({size:r="default",mobile:v=!0})=>E.jsxs(E.Fragment,{children:[E.jsx("img",{src:Kh,height:ti[r].desktop,style:{height:ti[r].desktop},alt:"NetBird Logo",className:wt(v&&"hidden md:block","group-hover:opacity-80 transition-all")}),v&&E.jsx("img",{src:Jh,width:ti[r].mobile,style:{width:ti[r].mobile},alt:"NetBird Logo",className:wt(v&&"md:hidden ml-4")})]});function Uf(){return E.jsxs("a",{href:"https://netbird.io?utm_source=netbird-proxy&utm_medium=web&utm_campaign=powered_by",target:"_blank",rel:"noopener noreferrer",className:"flex items-center justify-center mt-8 gap-2 group cursor-pointer",children:[E.jsx("span",{className:"text-sm text-nb-gray-400 font-light text-center group-hover:opacity-80 transition-all",children:"Powered by"}),E.jsx(kh,{size:"small",mobile:!1})]})}const Wh=({className:r})=>E.jsx("div",{className:wt("h-full w-full absolute left-0 top-0 rounded-md overflow-hidden z-0 pointer-events-none",r),children:E.jsx("div",{className:"bg-linear-to-b from-nb-gray-900/10 via-transparent to-transparent w-full h-full rounded-md"})}),Fd=({children:r,className:v})=>E.jsxs("div",{className:wt("px-6 sm:px-10 py-10 pt-8","bg-nb-gray-940 border border-nb-gray-910 rounded-lg relative",v),children:[E.jsx(Wh,{}),r]});function Cf({children:r,className:v}){return E.jsx("h1",{className:wt("text-xl! text-center z-10 relative",v),children:r})}function jf({children:r,className:v}){return E.jsx("div",{className:wt("text-sm text-nb-gray-300 font-light mt-2 block text-center z-10 relative",v),children:r})}const $h=()=>E.jsxs("div",{className:"flex items-center justify-center relative my-4",children:[E.jsx("span",{className:"bg-nb-gray-940 relative z-10 px-4 text-xs text-nb-gray-400 font-medium",children:"OR"}),E.jsx("span",{className:"h-px bg-nb-gray-900 w-full absolute z-0"})]}),Fh=({error:r})=>E.jsx("div",{className:"text-red-400 bg-red-800/20 border border-red-800/50 rounded-lg px-4 py-3 whitespace-break-spaces text-sm",children:r});function Id({className:r,htmlFor:v,...S}){return E.jsx("label",{htmlFor:v,className:wt("text-sm font-medium tracking-wider leading-none","peer-disabled:cursor-not-allowed peer-disabled:opacity-70","mb-2.5 inline-block text-nb-gray-200","flex items-center gap-2 select-none",r),...S})}const _f=t0(),It=_f.methods&&Object.keys(_f.methods).length>0?_f.methods:{password:"password",pin:"pin",oidc:"/auth/oidc"};function Ih(){xt.useEffect(()=>{document.title="Authentication Required - NetBird Service"},[]);const[r,v]=xt.useState(null),[S,f]=xt.useState(null),[_,O]=xt.useState(""),[D,U]=xt.useState(""),N=xt.useRef(null),p=xt.useRef(null),[R,H]=xt.useState(It.password?"password":"pin"),L=(nt,Ot)=>{v(Ot),f(null),nt==="password"?(U(""),setTimeout(()=>N.current?.focus(),200)):(O(""),setTimeout(()=>p.current?.focus(),200))},ot=(nt,Ot)=>{v(null),f(nt);const J=new FormData;nt==="password"?J.append(It.password,Ot):J.append(It.pin,Ot),fetch(globalThis.location.href,{method:"POST",body:J,redirect:"manual"}).then(Nt=>{if(Nt.type==="opaqueredirect"||Nt.status===0)f("redirect"),globalThis.location.reload();else if(Nt.status===429){const Xt=Number(Nt.headers.get("Retry-After")),pl=Number.isFinite(Xt)&&Xt>0?` Try again in ${Math.ceil(Xt)} seconds.`:" Please try again later.";L(nt,`Too many authentication attempts.${pl}`)}else L(nt,"Authentication failed. Please try again.")}).catch(()=>{L(nt,"An error occurred. Please try again.")})},ct=nt=>{O(nt),nt.length===6&&ot("pin",nt)},G=_.length===6,Q=D.length>0,V=S!==null||R==="password"&&!Q||R==="pin"&&!G,gt=It.password||It.pin,zt=It.password&&It.pin,_t=R==="password"?"Sign in":"Submit";return S==="redirect"?E.jsxs("main",{className:"mt-20",children:[E.jsxs(Fd,{className:"max-w-105 mx-auto",children:[E.jsx(Cf,{children:"Authenticated"}),E.jsx(jf,{children:"Loading service..."}),E.jsx("div",{className:"flex justify-center mt-7",children:E.jsx(kd,{className:"animate-spin",size:24})})]}),E.jsx(Uf,{})]}):E.jsxs("main",{className:"mt-20",children:[E.jsxs(Fd,{className:"max-w-105 mx-auto",children:[E.jsx(Cf,{children:"Authentication Required"}),E.jsx(jf,{children:"The service you are trying to access is protected. Please authenticate to continue."}),E.jsxs("div",{className:"flex flex-col gap-4 mt-7 z-10 relative",children:[r&&E.jsx(Fh,{error:r}),It.oidc&&E.jsxs(Ru,{variant:"primary",className:"w-full",onClick:()=>{globalThis.location.href=It.oidc},children:[E.jsx(Im,{size:16}),"Sign in with SSO"]}),It.oidc&>&&E.jsx($h,{}),gt&&E.jsxs("form",{onSubmit:nt=>{nt.preventDefault(),ot(R,R==="password"?D:_)},children:[zt&&E.jsx($e,{value:R,onChange:nt=>{H(nt),setTimeout(()=>{nt==="password"?N.current?.focus():p.current?.focus()},0)},children:E.jsxs($e.List,{className:"rounded-lg border mb-4",children:[E.jsxs($e.Trigger,{value:"password",children:[E.jsx(Fm,{size:14}),"Password"]}),E.jsxs($e.Trigger,{value:"pin",children:[E.jsx(Km,{size:14}),"PIN"]})]})}),E.jsxs("div",{className:"mb-4",children:[It.password&&(R==="password"||!It.pin)&&E.jsxs(E.Fragment,{children:[!zt&&E.jsx(Id,{htmlFor:"password",children:"Password"}),E.jsx(c0,{ref:N,type:"password",id:"password",placeholder:"Enter password",disabled:S!==null,showPasswordToggle:!0,autoFocus:!0,value:D,onChange:nt=>U(nt.target.value)})]}),It.pin&&(R==="pin"||!It.password)&&E.jsxs(E.Fragment,{children:[!zt&&E.jsx(Id,{htmlFor:"pin-0",children:"Enter PIN Code"}),E.jsx(Zh,{ref:p,value:_,onChange:ct,disabled:S!==null,autoFocus:!It.password})]})]}),E.jsx(Ru,{type:"submit",disabled:V,variant:"secondary",className:"w-full",children:S===null?_t:E.jsxs(E.Fragment,{children:[E.jsx(kd,{className:"animate-spin",size:16}),"Verifying..."]})})]})]})]}),E.jsx(Uf,{})]})}function Ph({success:r=!0}){return r?E.jsx("div",{className:"flex-1 flex items-center justify-center h-12 w-full px-5",children:E.jsx("div",{className:"w-full border-t-2 border-dashed border-green-500"})}):E.jsxs("div",{className:"flex-1 flex items-center justify-center h-12 min-w-10 px-5 relative",children:[E.jsx("div",{className:"w-full border-t-2 border-dashed border-nb-gray-900"}),E.jsx("div",{className:"absolute inset-0 flex items-center justify-center",children:E.jsx("div",{className:"w-8 h-8 rounded-full flex items-center justify-center",children:E.jsx(eh,{size:18,className:"text-netbird"})})})]})}function Of({icon:r,label:v,detail:S,success:f=!0,line:_=!0}){return E.jsxs(E.Fragment,{children:[_&&E.jsx(Ph,{success:f}),E.jsxs("div",{className:"flex flex-col items-center gap-2",children:[E.jsx("div",{className:"w-14 h-14 rounded-md flex items-center justify-center from-nb-gray-940 to-nb-gray-930/70 bg-gradient-to-br border border-nb-gray-910",children:E.jsx(r,{size:20,className:"text-nb-gray-200"})}),E.jsx("span",{className:"text-sm text-nb-gray-200 font-normal mt-1",children:v}),E.jsx("span",{className:`text-xs font-medium uppercase ${f?"text-green-500":"text-netbird"}`,children:f?"Connected":"Unreachable"}),S&&E.jsx("span",{className:"text-xs text-nb-gray-400 truncate text-center",children:S})]})]})}function tg({code:r,title:v,message:S,proxy:f=!0,destination:_=!0,requestId:O,simple:D=!1,retryUrl:U}){xt.useEffect(()=>{document.title=`${v} - NetBird Service`},[v]);const[N]=xt.useState(()=>new Date().toISOString());return E.jsxs("main",{className:"flex flex-col items-center mt-24 px-4 max-w-3xl mx-auto",children:[E.jsxs("div",{className:"text-sm text-netbird font-normal font-mono mb-3 z-10 relative",children:["Error ",r]}),E.jsx(Cf,{className:"text-3xl!",children:v}),E.jsx(jf,{className:"mt-2 mb-8 max-w-md",children:S}),!D&&E.jsxs("div",{className:"hidden sm:flex items-start justify-center w-full mt-6 mb-16 z-10 relative",children:[E.jsx(Of,{icon:th,label:"You",line:!1}),E.jsx(Of,{icon:lh,label:"Proxy",success:f}),E.jsx(Of,{icon:$m,label:"Destination",success:_})]}),E.jsxs("div",{className:"flex gap-3 justify-center items-center mb-6 z-10 relative",children:[E.jsxs(Ru,{variant:"primary",onClick:()=>{U?globalThis.location.href=U:globalThis.location.reload()},children:[E.jsx(Pm,{size:16}),"Refresh Page"]}),E.jsxs(Ru,{variant:"secondary",onClick:()=>globalThis.open("https://docs.netbird.io","_blank","noopener,noreferrer"),children:[E.jsx(Jm,{size:16}),"Documentation"]})]}),E.jsxs("div",{className:"text-center text-xs text-nb-gray-300 uppercase z-10 relative font-mono flex flex-col sm:flex-row gap-2 sm:gap-10 mt-4 mb-3",children:[E.jsxs("div",{children:[E.jsx("span",{className:"text-nb-gray-400",children:"REQUEST-ID:"})," ",O]}),E.jsxs("div",{children:[E.jsx("span",{className:"text-nb-gray-400",children:"TIMESTAMP:"})," ",N]})]}),E.jsx(Uf,{})]})}const Nf=t0();Zm.createRoot(document.getElementById("root")).render(E.jsx(xt.StrictMode,{children:Nf.page==="error"&&Nf.error?E.jsx(tg,{...Nf.error}):E.jsx(Ih,{})})); diff --git a/proxy/web/scripts/third-party-licenses.mjs b/proxy/web/scripts/third-party-licenses.mjs new file mode 100644 index 000000000..2ebff5f59 --- /dev/null +++ b/proxy/web/scripts/third-party-licenses.mjs @@ -0,0 +1,68 @@ +// Prints the third-party terms for the prebuilt UI in dist/ to stdout. +// Package versions come from package-lock.json and the license texts from an +// installed node_modules (run `npm ci --ignore-scripts` first). +import { existsSync, readdirSync, readFileSync } from "node:fs"; +import { dirname, join } from "node:path"; +import { fileURLToPath } from "node:url"; + +const webDir = join(dirname(fileURLToPath(import.meta.url)), ".."); + +// Build tools whose own code is emitted into dist: tailwindcss generates the +// preflight and utility CSS, vite injects its modulepreload polyfill. +const bundledTooling = new Set(["node_modules/tailwindcss", "node_modules/vite"]); + +// Inter ships as a font file rather than a package, so its OFL lives beside it. +const staticTerms = [ + { + title: "Inter 4.001 (git-66647c0bb), SIL Open Font License 1.1", + file: "src/assets/fonts/OFL.txt", + }, +]; + +const termPattern = /^(licen[cs]e|copying|notice|patents)/i; + +function fail(message) { + process.stderr.write(`${message}\n`); + process.exit(1); +} + +function shippedPackages(lock) { + return Object.entries(lock.packages) + .filter(([path, meta]) => path !== "" && (!meta.dev || bundledTooling.has(path))) + .sort(([a], [b]) => Number(a > b) - Number(a < b)); +} + +function packageSection(path, meta) { + const dir = join(webDir, path); + const manifest = join(dir, "package.json"); + if (!existsSync(manifest)) { + fail(`${path} is not installed; run npm ci --ignore-scripts in proxy/web`); + } + const installed = JSON.parse(readFileSync(manifest, "utf8")); + if (installed.version !== meta.version) { + fail(`${path} is ${installed.version}, package-lock.json pins ${meta.version}`); + } + + const terms = readdirSync(dir).filter((name) => termPattern.test(name)).sort(); + if (terms.length === 0) { + fail(`no license terms found for ${path}`); + } + + const name = path.slice(path.lastIndexOf("node_modules/") + "node_modules/".length); + return terms + .map((term) => `=== ${name} ${meta.version} (${term}) ===\n\n${readFileSync(join(dir, term), "utf8")}`) + .join("\n\n"); +} + +const lock = JSON.parse(readFileSync(join(webDir, "package-lock.json"), "utf8")); +const sections = shippedPackages(lock).map(([path, meta]) => packageSection(path, meta)); +for (const { title, file } of staticTerms) { + sections.push(`=== ${title} ===\n\n${readFileSync(join(webDir, file), "utf8")}`); +} + +process.stdout.write( + "Third-party terms for the prebuilt authentication UI in proxy/web/dist.\n" + + "Generated from proxy/web/package-lock.json at release time.\n\n\n" + + sections.join("\n\n\n") + + "\n", +); diff --git a/proxy/web/src/App.tsx b/proxy/web/src/App.tsx index ab453aa3e..3f09aacfb 100644 --- a/proxy/web/src/App.tsx +++ b/proxy/web/src/App.tsx @@ -68,6 +68,12 @@ function App() { if (res.type === "opaqueredirect" || res.status === 0) { setSubmitting("redirect"); globalThis.location.reload(); + } else if (res.status === 429) { + const seconds = Number(res.headers.get("Retry-After")); + const wait = Number.isFinite(seconds) && seconds > 0 + ? ` Try again in ${Math.ceil(seconds)} seconds.` + : " Please try again later."; + handleAuthError(method, `Too many authentication attempts.${wait}`); } else { handleAuthError(method, "Authentication failed. Please try again."); } diff --git a/proxy/web/src/assets/fonts/OFL.txt b/proxy/web/src/assets/fonts/OFL.txt new file mode 100644 index 000000000..9b2ca37b3 --- /dev/null +++ b/proxy/web/src/assets/fonts/OFL.txt @@ -0,0 +1,92 @@ +Copyright (c) 2016 The Inter Project Authors (https://github.com/rsms/inter) + +This Font Software is licensed under the SIL Open Font License, Version 1.1. +This license is copied below, and is also available with a FAQ at: +http://scripts.sil.org/OFL + +----------------------------------------------------------- +SIL OPEN FONT LICENSE Version 1.1 - 26 February 2007 +----------------------------------------------------------- + +PREAMBLE +The goals of the Open Font License (OFL) are to stimulate worldwide +development of collaborative font projects, to support the font creation +efforts of academic and linguistic communities, and to provide a free and +open framework in which fonts may be shared and improved in partnership +with others. + +The OFL allows the licensed fonts to be used, studied, modified and +redistributed freely as long as they are not sold by themselves. The +fonts, including any derivative works, can be bundled, embedded, +redistributed and/or sold with any software provided that any reserved +names are not used by derivative works. The fonts and derivatives, +however, cannot be released under any other type of license. The +requirement for fonts to remain under this license does not apply +to any document created using the fonts or their derivatives. + +DEFINITIONS +"Font Software" refers to the set of files released by the Copyright +Holder(s) under this license and clearly marked as such. This may +include source files, build scripts and documentation. + +"Reserved Font Name" refers to any names specified as such after the +copyright statement(s). + +"Original Version" refers to the collection of Font Software components as +distributed by the Copyright Holder(s). + +"Modified Version" refers to any derivative made by adding to, deleting, +or substituting -- in part or in whole -- any of the components of the +Original Version, by changing formats or by porting the Font Software to a +new environment. + +"Author" refers to any designer, engineer, programmer, technical +writer or other person who contributed to the Font Software. + +PERMISSION AND CONDITIONS +Permission is hereby granted, free of charge, to any person obtaining +a copy of the Font Software, to use, study, copy, merge, embed, modify, +redistribute, and sell modified and unmodified copies of the Font +Software, subject to the following conditions: + +1) Neither the Font Software nor any of its individual components, +in Original or Modified Versions, may be sold by itself. + +2) Original or Modified Versions of the Font Software may be bundled, +redistributed and/or sold with any software, provided that each copy +contains the above copyright notice and this license. These can be +included either as stand-alone text files, human-readable headers or +in the appropriate machine-readable metadata fields within text or +binary files as long as those fields can be easily viewed by the user. + +3) No Modified Version of the Font Software may use the Reserved Font +Name(s) unless explicit written permission is granted by the corresponding +Copyright Holder. This restriction only applies to the primary font name as +presented to the users. + +4) The name(s) of the Copyright Holder(s) or the Author(s) of the Font +Software shall not be used to promote, endorse or advertise any +Modified Version, except to acknowledge the contribution(s) of the +Copyright Holder(s) and the Author(s) or with their explicit written +permission. + +5) The Font Software, modified or unmodified, in part or in whole, +must be distributed entirely under this license, and must not be +distributed under any other license. The requirement for fonts to +remain under this license does not apply to any document created +using the Font Software. + +TERMINATION +This license becomes null and void if any of the above conditions are +not met. + +DISCLAIMER +THE FONT SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, +EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO ANY WARRANTIES OF +MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT +OF COPYRIGHT, PATENT, TRADEMARK, OR OTHER RIGHT. IN NO EVENT SHALL THE +COPYRIGHT HOLDER BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, +INCLUDING ANY GENERAL, SPECIAL, INDIRECT, INCIDENTAL, OR CONSEQUENTIAL +DAMAGES, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +FROM, OUT OF THE USE OR INABILITY TO USE THE FONT SOFTWARE OR FROM +OTHER DEALINGS IN THE FONT SOFTWARE. diff --git a/proxy/web/web.go b/proxy/web/web.go index 6773a9c1a..de3e4771a 100644 --- a/proxy/web/web.go +++ b/proxy/web/web.go @@ -10,6 +10,8 @@ import ( "net/url" "path/filepath" "strings" + + "github.com/netbirdio/netbird/proxy/auth" ) // PathPrefix is the unique URL prefix for serving the proxy's own web assets. @@ -180,7 +182,8 @@ func ServeAccessDeniedPage(w http.ResponseWriter, r *http.Request, code int, tit // stripAuthParams returns the request URI with auth-related query parameters removed. func stripAuthParams(u *url.URL) string { q := u.Query() - q.Del("session_token") + q.Del(auth.SessionTokenQueryParam) + q.Del(auth.SessionCodeQueryParam) q.Del("error") q.Del("error_description") clean := *u diff --git a/release_files/rpm-changelog.sh b/release_files/rpm-changelog.sh index 20af2d415..d9150399e 100755 --- a/release_files/rpm-changelog.sh +++ b/release_files/rpm-changelog.sh @@ -7,6 +7,12 @@ # template headings, review checklists, HTML comments and Co-authored-by # trailers. None of that belongs in a package on Red Hat's catalog, and it is # most of the changelog's size. Keep the subject line and drop the rest. +# +# chglog also records the bare tag as each entry's version, while nfpm writes +# that string into the changelog header verbatim and never appends the release. +# rpmlint then reports incoherent-version-in-changelog, because the entry reads +# 0.79.0 while the package is 0.79.0-1. Rewrite each version the way nfpm +# renders the package EVR. set -eu @@ -20,14 +26,29 @@ path = sys.argv[1] lines = open(path, encoding="utf-8").read().split("\n") NOTE = re.compile(r"^ note: (.*)$") +SEMVER = re.compile(r"^- semver: (.*)$") BLOCK = {"|", "|-", "|+", ">", ">-", ">+"} +# nfpm defaults the RPM release to 1 and the packaging sets no other value. +RELEASE = "1" + def quote(text): """Render text as a YAML single-quoted scalar.""" return " note: '{}'".format(text.replace("'", "''")) +def evr(version): + """Render a semver tag the way nfpm renders the package EVR.""" + version, _, metadata = version.partition("+") + core, _, prerelease = version.partition("-") + if prerelease: + core += "~" + prerelease.replace("-", "_") + if metadata: + core += "+" + metadata + return "{}-{}".format(core, RELEASE) + + def first_line_of_double_quoted(value): """Text of a double-quoted scalar up to its first \\n escape.""" out = [] @@ -52,6 +73,13 @@ seen = 0 i = 0 while i < len(lines): line = lines[i] + + m = SEMVER.match(line) + if m: + out.append("- semver: '{}'".format(evr(m.group(1)))) + i += 1 + continue + m = NOTE.match(line) if not m: out.append(line) @@ -107,4 +135,11 @@ if grep -nE '^ note: ".*\\n' changelog.yml; then exit 1 fi +# Every entry must carry the release, or rpmlint reports the changelog version +# as incoherent with the package again. +if grep -nE "^- semver: " changelog.yml | grep -vE -- "-[0-9]+'$"; then + echo "changelog entries without the RPM release survived the rewrite" >&2 + exit 1 +fi + test -s changelog.yml diff --git a/release_files/rpm-provides.sh b/release_files/rpm-provides.sh new file mode 100644 index 000000000..b1332c18b --- /dev/null +++ b/release_files/rpm-provides.sh @@ -0,0 +1,46 @@ +#!/bin/sh +# +# Write .goreleaser.generated.yaml with the @RPM_EVR@ placeholder filled in. +# +# Red Hat certification (RPM Version Handling) expects rpmbuild's ISA provide, +# netbird(x86-64) = . nfpm does not emit it and GoReleaser does not template +# the provides field, so the version is substituted before GoReleaser runs. +# +# The value has to match what nfpm derives from the same tag: a semver +# prerelease becomes a tilde suffix, and the release defaults to 1. + +set -eu + +OUT=.goreleaser.generated.yaml + +TAG="${GITHUB_REF#refs/tags/}" +case "$TAG" in +v*) ;; +*) TAG=$(git describe --tags --abbrev=0) ;; +esac + +EVR=$(python3 - "$TAG" <<'PYEOF' +import sys + +version = sys.argv[1].lstrip("v") +version, _, metadata = version.partition("+") +core, _, prerelease = version.partition("-") +if prerelease: + core += "~" + prerelease.replace("-", "_") +if metadata: + core += "+" + metadata +print("{}-1".format(core)) +PYEOF +) + +# Written to a separate, ignored file: GoReleaser refuses to release from a +# dirty tree, so .goreleaser.yaml itself must stay untouched. +sed "s/@RPM_EVR@/${EVR}/g" .goreleaser.yaml > "$OUT" + +# A surviving placeholder means the provides entries moved or were renamed. +if grep -n "@RPM_EVR@" "$OUT"; then + echo "unsubstituted @RPM_EVR@ left in $OUT" >&2 + exit 1 +fi + +echo "rpm provides version: ${EVR} -> ${OUT}" diff --git a/route/route.go b/route/route.go index 3bdb0a3a1..ef9a39ef7 100644 --- a/route/route.go +++ b/route/route.go @@ -95,7 +95,7 @@ type Route struct { ID ID `gorm:"primaryKey"` // AccountID is a reference to Account that this object belongs AccountID string `gorm:"index"` - PublicID string `json:"-"` + PublicID string `json:"-" gorm:"index"` // Network and Domains are mutually exclusive Network netip.Prefix `gorm:"serializer:json"` Domains domain.List `gorm:"serializer:json"` diff --git a/shared/auth/jwt/validator.go b/shared/auth/jwt/validator.go index 62e127751..240d64bb3 100644 --- a/shared/auth/jwt/validator.go +++ b/shared/auth/jwt/validator.go @@ -18,7 +18,6 @@ import ( "time" "github.com/golang-jwt/jwt/v5" - log "github.com/sirupsen/logrus" ) @@ -217,8 +216,8 @@ func (v *Validator) ValidateAndParse(ctx context.Context, token string) (*jwt.To jwt.WithAudience(v.audienceList...), jwt.WithIssuer(v.issuer), jwt.WithIssuedAt(), + jwt.WithStrictDecoding(), ) - // Check if there was an error in parsing... if err != nil { err = fmt.Errorf("%w: %s", errTokenParsing, err) diff --git a/shared/lifecycle/stop_handlers.go b/shared/lifecycle/stop_handlers.go new file mode 100644 index 000000000..f6ec2688b --- /dev/null +++ b/shared/lifecycle/stop_handlers.go @@ -0,0 +1,57 @@ +package lifecycle + +import ( + "runtime/debug" + "sync" + + log "github.com/sirupsen/logrus" +) + +// StopHandlers collects functions to run once when their owner exits. Embed it +// in a server type to expose OnStop and RunStopHandlers. +type StopHandlers struct { + mu sync.Mutex + stopped bool + handlers []func() +} + +// OnStop registers fn to run once when the owner stops. Handlers run in +// reverse registration order. A handler registered after the owner has +// stopped runs immediately. +func (h *StopHandlers) OnStop(fn func()) { + h.mu.Lock() + stopped := h.stopped + if !stopped { + h.handlers = append(h.handlers, fn) + } + h.mu.Unlock() + + if stopped { + runStopHandler(fn) + } +} + +// RunStopHandlers runs every registered handler once, last registered first. +// Later calls are no-ops, so it can be wired to several exit paths at once. +func (h *StopHandlers) RunStopHandlers() { + h.mu.Lock() + handlers := h.handlers + h.handlers = nil + h.stopped = true + h.mu.Unlock() + + for i := len(handlers) - 1; i >= 0; i-- { + runStopHandler(handlers[i]) + } +} + +// runStopHandler keeps one panicking handler from skipping the ones still +// pending; on the shutdown path there is no second chance to run them. +func runStopHandler(fn func()) { + defer func() { + if r := recover(); r != nil { + log.Errorf("stop handler panicked: %v\n%s", r, debug.Stack()) + } + }() + fn() +} diff --git a/shared/lifecycle/stop_handlers_test.go b/shared/lifecycle/stop_handlers_test.go new file mode 100644 index 000000000..787f39e6d --- /dev/null +++ b/shared/lifecycle/stop_handlers_test.go @@ -0,0 +1,43 @@ +package lifecycle + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestStopHandlers_RunOnceInReverseOrder(t *testing.T) { + var h StopHandlers + var order []string + h.OnStop(func() { order = append(order, "first") }) + h.OnStop(func() { order = append(order, "second") }) + + h.RunStopHandlers() + h.RunStopHandlers() + + assert.Equal(t, []string{"second", "first"}, order, "handlers must run once, last registered first") +} + +func TestStopHandlers_PanicDoesNotSkipRemainingHandlers(t *testing.T) { + var h StopHandlers + var order []string + h.OnStop(func() { order = append(order, "first") }) + h.OnStop(func() { panic("boom") }) + h.OnStop(func() { order = append(order, "third") }) + + h.RunStopHandlers() + + assert.Equal(t, []string{"third", "first"}, order, "handlers around a panicking one must still run") +} + +func TestStopHandlers_LateRegistrationRunsImmediately(t *testing.T) { + var h StopHandlers + h.RunStopHandlers() + + runs := 0 + h.OnStop(func() { runs++ }) + assert.Equal(t, 1, runs, "a handler registered after the stop must run right away") + + h.RunStopHandlers() + assert.Equal(t, 1, runs, "later runs must stay no-ops and must not repeat the handler") +} diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 883270597..fe030b633 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -826,6 +826,20 @@ components: - ssh_enabled - login_expiration_enabled - inactivity_expiration_enabled + NetworkAddress: + type: object + properties: + net_ip: + description: IP address with CIDR of the interface + type: string + example: 192.168.0.11/24 + mac: + description: MAC address of the interface + type: string + example: "00:93:37:bd:83:0f" + required: + - net_ip + - mac Peer: allOf: - $ref: '#/components/schemas/PeerMinimum' @@ -845,6 +859,11 @@ components: type: string format: ipv6 example: "fd00:4e42:ab12::1" + network_addresses: + description: Network interfaces (IP + MAC) reported by the peer + type: array + items: + $ref: '#/components/schemas/NetworkAddress' connection_ip: description: Peer's public connection IP address type: string @@ -6467,8 +6486,8 @@ components: example: "d1m3kebd9pcs0c1pnu7g" state: type: string - description: Derived deployment state. `provisioning` until the gateway is rolled out and connected, `ready` while the gateway actively serves the endpoint, `failed` when the rollout reported a failure. - enum: [ "provisioning", "ready", "failed" ] + description: Derived deployment state. `provisioning` until the gateway is rolled out and connected, `ready` while the gateway actively serves the endpoint, `failed` when the rollout reported a failure, `disabled` while the gateway is turned off and the endpoint is not served. + enum: [ "provisioning", "ready", "failed", "disabled" ] example: "ready" endpoint: type: string @@ -7530,6 +7549,11 @@ paths: schema: type: string description: Filter peers by IP address + - in: query + name: mac + schema: + type: string + description: Filter peers by MAC address of a network interface security: - BearerAuth: [ ] - TokenAuth: [ ] @@ -10615,6 +10639,28 @@ paths: $ref: "#/components/responses/requires_authentication" "500": $ref: "#/components/responses/internal_error" + delete: + summary: Delete MSP tenant + tags: + - MSP + parameters: + - in: path + name: id + required: true + schema: + type: string + description: The unique identifier of a tenant account + responses: + "200": + description: Successfully deleted the tenant + "400": + $ref: "#/components/responses/bad_request" + "403": + $ref: "#/components/responses/requires_authentication" + "404": + description: The tenant was not found + "500": + $ref: "#/components/responses/internal_error" /api/integrations/msp/tenants/{id}/unlink: post: summary: Unlink a tenant @@ -13361,7 +13407,7 @@ paths: /api/reverse-proxies/domains/{domainId}: delete: summary: Delete a Custom domain - description: Delete an existing service custom domain + description: Delete an existing service custom domain after removing or moving all services that use it or its subdomains, including disabled services. tags: [ Services ] security: - BearerAuth: [ ] @@ -13384,6 +13430,9 @@ paths: "$ref": "#/components/responses/forbidden" '404': "$ref": "#/components/responses/not_found" + '412': + description: The domain or one of its subdomains is still used by a service + content: { } '500': "$ref": "#/components/responses/internal_error" /api/reverse-proxies/domains/{domainId}/validate: @@ -13586,7 +13635,7 @@ paths: /api/integrations/agent-network/managed-proxy: post: summary: Provision a managed Agent Network gateway - description: Starts provisioning of a NetBird-managed Agent Network gateway for the account, allocating its endpoint under the managed zone on the first call. Idempotent — answers 202 when this call started (or, after a failure, restarted) provisioning and 200 when a deployment already exists, reporting current state either way. Returns 409 when the account already has an Agent Network endpoint that managed provisioning does not own, and 503 when endpoint allocation is temporarily exhausted. + description: Starts provisioning of a NetBird-managed Agent Network gateway for the account, allocating its endpoint under the managed zone on the first call. Idempotent — answers 202 when this call started (or, after a failure, restarted) provisioning and 200 when a deployment already exists, reporting current state either way. A disabled deployment answers 200 with state `disabled` and stays disabled. Returns 409 when the account already has an Agent Network endpoint that managed provisioning does not own, and 503 when endpoint allocation is temporarily exhausted. tags: [ Agent Network ] security: - BearerAuth: [ ] @@ -13624,7 +13673,7 @@ paths: $ref: '#/components/schemas/ErrorResponse' get: summary: Retrieve managed Agent Network gateway status - description: Reports the account's managed gateway deployment and its derived state. Returns 404 when the account has no managed deployment. + description: Reports the account's managed gateway deployment and its derived state, including `disabled` for a deployment that is turned off. Returns 404 when the account has no managed deployment. tags: [ Agent Network ] security: - BearerAuth: [ ] diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index e47e53b9b..20b6f4701 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -79,6 +79,7 @@ func (e AgentNetworkConsumptionDimensionKind) Valid() bool { // Defines values for AgentNetworkManagedProxyState. const ( + AgentNetworkManagedProxyStateDisabled AgentNetworkManagedProxyState = "disabled" AgentNetworkManagedProxyStateFailed AgentNetworkManagedProxyState = "failed" AgentNetworkManagedProxyStateProvisioning AgentNetworkManagedProxyState = "provisioning" AgentNetworkManagedProxyStateReady AgentNetworkManagedProxyState = "ready" @@ -87,6 +88,8 @@ const ( // Valid indicates whether the value is a known member of the AgentNetworkManagedProxyState enum. func (e AgentNetworkManagedProxyState) Valid() bool { switch e { + case AgentNetworkManagedProxyStateDisabled: + return true case AgentNetworkManagedProxyStateFailed: return true case AgentNetworkManagedProxyStateProvisioning: @@ -2259,11 +2262,11 @@ type AgentNetworkManagedProxy struct { // Region Region of the cluster hosting the deployment. Region *string `json:"region,omitempty"` - // State Derived deployment state. `provisioning` until the gateway is rolled out and connected, `ready` while the gateway actively serves the endpoint, `failed` when the rollout reported a failure. + // State Derived deployment state. `provisioning` until the gateway is rolled out and connected, `ready` while the gateway actively serves the endpoint, `failed` when the rollout reported a failure, `disabled` while the gateway is turned off and the endpoint is not served. State AgentNetworkManagedProxyState `json:"state"` } -// AgentNetworkManagedProxyState Derived deployment state. `provisioning` until the gateway is rolled out and connected, `ready` while the gateway actively serves the endpoint, `failed` when the rollout reported a failure. +// AgentNetworkManagedProxyState Derived deployment state. `provisioning` until the gateway is rolled out and connected, `ready` while the gateway actively serves the endpoint, `failed` when the rollout reported a failure, `disabled` while the gateway is turned off and the endpoint is not served. type AgentNetworkManagedProxyState string // AgentNetworkManagedProxyConflict Conflict body returned when the account already has an Agent Network endpoint that managed provisioning does not own, naming that endpoint. @@ -3835,6 +3838,15 @@ type Network struct { RoutingPeersCount int `json:"routing_peers_count"` } +// NetworkAddress defines model for NetworkAddress. +type NetworkAddress struct { + // Mac MAC address of the interface + Mac string `json:"mac"` + + // NetIp IP address with CIDR of the interface + NetIp string `json:"net_ip"` +} + // NetworkRequest defines model for NetworkRequest. type NetworkRequest struct { // Description Network description @@ -4284,6 +4296,9 @@ type Peer struct { // Name Peer's hostname Name string `json:"name"` + // NetworkAddresses Network interfaces (IP + MAC) reported by the peer + NetworkAddresses *[]NetworkAddress `json:"network_addresses,omitempty"` + // Os Peer's operating system and version Os string `json:"os"` @@ -4378,6 +4393,9 @@ type PeerBatch struct { // Name Peer's hostname Name string `json:"name"` + // NetworkAddresses Network interfaces (IP + MAC) reported by the peer + NetworkAddresses *[]NetworkAddress `json:"network_addresses,omitempty"` + // Os Peer's operating system and version Os string `json:"os"` @@ -6300,6 +6318,9 @@ type GetApiPeersParams struct { // Ip Filter peers by IP address Ip *string `form:"ip,omitempty" json:"ip,omitempty"` + + // Mac Filter peers by MAC address of a network interface + Mac *string `form:"mac,omitempty" json:"mac,omitempty"` } // GetApiPeersPeerIdIngressPortsParams defines parameters for GetApiPeersPeerIdIngressPorts. diff --git a/shared/management/networkmap/envelope.go b/shared/management/networkmap/envelope.go index e7961fd7b..fd9dd6bbd 100644 --- a/shared/management/networkmap/envelope.go +++ b/shared/management/networkmap/envelope.go @@ -35,7 +35,12 @@ type EnvelopeResult struct { // // dnsName is the account's DNS domain ("netbird.cloud" etc.); used when // rebuilding the per-peer FQDNs that proto.RemotePeerConfig carries. -func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, localPeerKey, dnsName string) (*EnvelopeResult, error) { +// +// skipRouteFirewallRules leaves RoutesFirewallRules empty. Callers that have +// no firewall to program pass true: the rules are the most expensive part of +// Calculate on a peer that routes many network resources, and nothing reads +// them afterwards. +func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, localPeerKey, dnsName string, skipRouteFirewallRules bool) (*EnvelopeResult, error) { components, err := DecodeEnvelope(ctx, env) if err != nil { return nil, fmt.Errorf("decode envelope: %w", err) @@ -53,6 +58,7 @@ func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, lo return nil, fmt.Errorf("receiving peer (wg_key prefix %q) not found among %d decoded peers — components have no PeerID, Calculate would return empty", trimKey(localPeerKey), len(components.Peers)) } components.PeerID = canonicalKey + components.SkipRouteFirewallRules = skipRouteFirewallRules includeIPv6 := localPeer.SupportsIPv6() && localPeer.IPv6.IsValid() useSourcePrefixes := localPeer.SupportsSourcePrefixes() diff --git a/shared/management/networkmap/envelope_test.go b/shared/management/networkmap/envelope_test.go index 7fe2a5277..98333a3c8 100644 --- a/shared/management/networkmap/envelope_test.go +++ b/shared/management/networkmap/envelope_test.go @@ -9,6 +9,7 @@ import ( "net/netip" "testing" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" goproto "google.golang.org/protobuf/proto" @@ -37,7 +38,7 @@ func TestEnvelopeToNetworkMap_RoundTrip(t *testing.T) { var decoded proto.NetworkMapEnvelope require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") - result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false) require.NoError(t, err, "EnvelopeToNetworkMap") require.NotNil(t, result) require.NotNil(t, result.NetworkMap, "decoded NetworkMap must be non-nil") @@ -78,7 +79,7 @@ func TestCalculate_FirewallRuleProtocol_NeverNetbirdSSH(t *testing.T) { var decoded proto.NetworkMapEnvelope require.NoError(t, goproto.Unmarshal(wire, &decoded)) - result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false) require.NoError(t, err) require.NotEmpty(t, result.NetworkMap.FirewallRules, "ssh policy should produce firewall rules") for i, fr := range result.NetworkMap.FirewallRules { @@ -88,13 +89,13 @@ func TestCalculate_FirewallRuleProtocol_NeverNetbirdSSH(t *testing.T) { } func TestEnvelopeToNetworkMap_NilEnvelope(t *testing.T) { - _, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), nil, "key", "netbird.cloud") + _, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), nil, "key", "netbird.cloud", false) require.Error(t, err, "nil envelope must produce an error rather than panic") } func TestEnvelopeToNetworkMap_FullPayloadMissing(t *testing.T) { env := &proto.NetworkMapEnvelope{} - _, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), env, "key", "netbird.cloud") + _, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), env, "key", "netbird.cloud", false) require.Error(t, err, "envelope with no Full payload must produce an error") } @@ -126,7 +127,7 @@ func TestDecodeEnvelope_MalformedWgKeyPeerSkipped(t *testing.T) { var decoded proto.NetworkMapEnvelope require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") - result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false) require.NoError(t, err, "EnvelopeToNetworkMap must tolerate one bad peer key") require.NotNil(t, result) require.NotNil(t, result.Components) @@ -195,7 +196,7 @@ func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) { var decodedEnv proto.NetworkMapEnvelope require.NoError(t, goproto.Unmarshal(wire, &decodedEnv), "unmarshal envelope") - result, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decodedEnv, peers["peer-T"].Key, "netbird.cloud") + result, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decodedEnv, peers["peer-T"].Key, "netbird.cloud", false) require.NoError(t, err, "EnvelopeToNetworkMap") clientNM := result.NetworkMap @@ -253,7 +254,7 @@ func TestEnvelopeToNetworkMap_EmptyComponents(t *testing.T) { var decoded proto.NetworkMapEnvelope require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") - result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false) require.NoError(t, err, "EnvelopeToNetworkMap must degrade gracefully on empty components") require.Equal(t, uint64(7), result.NetworkMap.Serial) require.Empty(t, result.NetworkMap.RemotePeers, "unvalidated peer connects to nobody") @@ -276,7 +277,7 @@ func TestEnvelopeToNetworkMap_MissingNetwork(t *testing.T) { var decoded proto.NetworkMapEnvelope require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") - result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud", false) require.NoError(t, err, "a missing AccountNetwork must not panic the client") require.NotNil(t, result.Components.Network) require.NotEmpty(t, result.NetworkMap.RemotePeers, "the rest of the snapshot stays usable") @@ -353,3 +354,110 @@ func randomWgKey(t *testing.T) string { require.NoError(t, err) return base64.StdEncoding.EncodeToString(raw[:]) } + +// TestEnvelopeToNetworkMap_SkipRouteFirewallRules covers the flag end to end, +// through the envelope rather than by poking Calculate directly. The +// RoutesFirewallRulesIsEmpty derivation is the part that matters: the client's +// legacy-management probe reads an empty rule list together with that bit, so +// skipping the rules must set it rather than leave it false. +func TestEnvelopeToNetworkMap_SkipRouteFirewallRules(t *testing.T) { + ctx := context.Background() + c, routerKey := buildRoutedResourceComponents(t) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + full, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decoded, routerKey, "netbird.cloud", false) + require.NoError(t, err, "EnvelopeToNetworkMap without skip") + require.NotEmpty(t, full.NetworkMap.RoutesFirewallRules, + "baseline: the router peer must receive route firewall rules") + require.False(t, full.NetworkMap.RoutesFirewallRulesIsEmpty, + "baseline: the empty bit must be false when rules are present") + + var decodedSkip proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decodedSkip), "unmarshal envelope") + skipped, err := nbnetworkmap.EnvelopeToNetworkMap(ctx, &decodedSkip, routerKey, "netbird.cloud", true) + require.NoError(t, err, "EnvelopeToNetworkMap with skip") + + assert.Empty(t, skipped.NetworkMap.RoutesFirewallRules, + "route firewall rules must not be computed when skipped") + assert.True(t, skipped.NetworkMap.RoutesFirewallRulesIsEmpty, + "the empty bit must be derived from the skipped list, or the client misreads it as legacy management") + assert.Len(t, skipped.NetworkMap.Routes, len(full.NetworkMap.Routes), + "skipping route firewall rules must not change the routes") + assert.Len(t, skipped.NetworkMap.RemotePeers, len(full.NetworkMap.RemotePeers), + "skipping route firewall rules must not change the remote peers") +} + +// buildRoutedResourceComponents returns components in which the local peer is +// the routing peer for one enabled network resource, reachable by a second +// peer through a resource policy — the minimum shape that yields a non-empty +// RoutesFirewallRules. It also returns the local peer's WG key. +func buildRoutedResourceComponents(t *testing.T) (*types.NetworkMapComponents, string) { + t.Helper() + + routerKey := randomWgKey(t) + peers := map[string]*nmdata.Peer{ + "peer-R": { + ID: "peer-R", Key: routerKey, DNSLabel: "router", + IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}), + Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"}, + }, + "peer-S": { + ID: "peer-S", Key: randomWgKey(t), DNSLabel: "source", + IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}), + Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"}, + }, + } + + resourcePolicy := &nmdata.Policy{ + ID: "pol-res", PublicID: "10", Enabled: true, + Rules: []*nmdata.PolicyRule{{ + ID: "rule-res", + Enabled: true, + Action: string(types.PolicyTrafficActionAccept), + Protocol: string(types.PolicyRuleProtocolALL), + Sources: []string{"g-src"}, + }}, + } + + c := &types.NetworkMapComponents{ + PeerID: "peer-R", + Network: &nmdata.Network{ + Identifier: "net-routed-resource", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 1, + }, + AccountSettings: &nmdata.AccountSettingsInfo{}, + DNSSettings: &nmdata.DNSSettings{}, + Peers: peers, + Groups: map[string]*nmdata.Group{ + "g-src": {PublicID: "1", Name: "sources", Peers: []string{"peer-S"}}, + "g-routers": {PublicID: "2", Name: "routers", Peers: []string{"peer-R"}}, + }, + NetworkResources: []*nmdata.NetworkResource{{ + ID: "res-1", NetworkID: "netid-1", PublicID: "100", Name: "res1", + Type: "subnet", + Prefix: netip.MustParsePrefix("10.200.0.0/24"), + Enabled: true, + }}, + RoutersMap: map[string]map[string]*nmdata.NetworkRouter{ + "netid-1": {"peer-R": { + PublicID: "200", PeerGroups: []string{"g-routers"}, Metric: 9999, Enabled: true, + }}, + }, + ResourcePoliciesMap: map[string][]*nmdata.Policy{ + "res-1": {resourcePolicy}, + }, + Policies: []*nmdata.Policy{resourcePolicy}, + NetworkXIDToPublicID: map[string]string{"netid-1": "1"}, + } + + return c, routerKey +} diff --git a/shared/management/proto/proxy_service.pb.go b/shared/management/proto/proxy_service.pb.go index 496774a4b..09dff0d36 100644 --- a/shared/management/proto/proxy_service.pb.go +++ b/shared/management/proto/proxy_service.pb.go @@ -2265,8 +2265,11 @@ type ValidateSessionRequest struct { sizeCache protoimpl.SizeCache unknownFields protoimpl.UnknownFields - Domain string `protobuf:"bytes,1,opt,name=domain,proto3" json:"domain,omitempty"` + Domain string `protobuf:"bytes,1,opt,name=domain,proto3" json:"domain,omitempty"` + // Deprecated: Do not use. SessionToken string `protobuf:"bytes,2,opt,name=session_token,json=sessionToken,proto3" json:"session_token,omitempty"` + // session_code is a short-lived, single-use code exchanged for a session token. + SessionCode string `protobuf:"bytes,3,opt,name=session_code,json=sessionCode,proto3" json:"session_code,omitempty"` } func (x *ValidateSessionRequest) Reset() { @@ -2308,6 +2311,7 @@ func (x *ValidateSessionRequest) GetDomain() string { return "" } +// Deprecated: Do not use. func (x *ValidateSessionRequest) GetSessionToken() string { if x != nil { return x.SessionToken @@ -2315,6 +2319,13 @@ func (x *ValidateSessionRequest) GetSessionToken() string { return "" } +func (x *ValidateSessionRequest) GetSessionCode() string { + if x != nil { + return x.SessionCode + } + return "" +} + type ValidateSessionResponse struct { state protoimpl.MessageState sizeCache protoimpl.SizeCache @@ -2333,6 +2344,8 @@ type ValidateSessionResponse struct { // Stamped onto upstream requests as X-NetBird-Groups so downstream // services can read names rather than opaque ids. PeerGroupNames []string `protobuf:"bytes,6,rep,name=peer_group_names,json=peerGroupNames,proto3" json:"peer_group_names,omitempty"` + // session_token contains the durable token issued when session_code is redeemed. + SessionToken string `protobuf:"bytes,7,opt,name=session_token,json=sessionToken,proto3" json:"session_token,omitempty"` } func (x *ValidateSessionResponse) Reset() { @@ -2409,6 +2422,13 @@ func (x *ValidateSessionResponse) GetPeerGroupNames() []string { return nil } +func (x *ValidateSessionResponse) GetSessionToken() string { + if x != nil { + return x.SessionToken + } + return "" +} + // ValidateTunnelPeerRequest carries the inbound peer's tunnel IP and the // service domain whose group requirements should gate access. The calling // account is inferred from the proxy's gRPC metadata (ProxyToken). @@ -3519,223 +3539,228 @@ var file_proxy_service_proto_rawDesc = []byte{ 0x09, 0x52, 0x0b, 0x72, 0x65, 0x64, 0x69, 0x72, 0x65, 0x63, 0x74, 0x55, 0x72, 0x6c, 0x22, 0x26, 0x0a, 0x12, 0x47, 0x65, 0x74, 0x4f, 0x49, 0x44, 0x43, 0x55, 0x52, 0x4c, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x10, 0x0a, 0x03, 0x75, 0x72, 0x6c, 0x18, 0x01, 0x20, 0x01, 0x28, - 0x09, 0x52, 0x03, 0x75, 0x72, 0x6c, 0x22, 0x55, 0x0a, 0x16, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, + 0x09, 0x52, 0x03, 0x75, 0x72, 0x6c, 0x22, 0x7c, 0x0a, 0x16, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x16, 0x0a, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, - 0x52, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, 0x12, 0x23, 0x0a, 0x0d, 0x73, 0x65, 0x73, 0x73, - 0x69, 0x6f, 0x6e, 0x5f, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, - 0x0c, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x54, 0x6f, 0x6b, 0x65, 0x6e, 0x22, 0xdc, 0x01, - 0x0a, 0x17, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, 0x73, 0x69, 0x6f, - 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x14, 0x0a, 0x05, 0x76, 0x61, 0x6c, + 0x52, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, 0x12, 0x27, 0x0a, 0x0d, 0x73, 0x65, 0x73, 0x73, + 0x69, 0x6f, 0x6e, 0x5f, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x42, + 0x02, 0x18, 0x01, 0x52, 0x0c, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x54, 0x6f, 0x6b, 0x65, + 0x6e, 0x12, 0x21, 0x0a, 0x0c, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x5f, 0x63, 0x6f, 0x64, + 0x65, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0b, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, + 0x43, 0x6f, 0x64, 0x65, 0x22, 0x81, 0x02, 0x0a, 0x17, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, + 0x65, 0x53, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, + 0x12, 0x14, 0x0a, 0x05, 0x76, 0x61, 0x6c, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x08, 0x52, + 0x05, 0x76, 0x61, 0x6c, 0x69, 0x64, 0x12, 0x17, 0x0a, 0x07, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x69, + 0x64, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, 0x75, 0x73, 0x65, 0x72, 0x49, 0x64, 0x12, + 0x1d, 0x0a, 0x0a, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x65, 0x6d, 0x61, 0x69, 0x6c, 0x18, 0x03, 0x20, + 0x01, 0x28, 0x09, 0x52, 0x09, 0x75, 0x73, 0x65, 0x72, 0x45, 0x6d, 0x61, 0x69, 0x6c, 0x12, 0x23, + 0x0a, 0x0d, 0x64, 0x65, 0x6e, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x18, + 0x04, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0c, 0x64, 0x65, 0x6e, 0x69, 0x65, 0x64, 0x52, 0x65, 0x61, + 0x73, 0x6f, 0x6e, 0x12, 0x24, 0x0a, 0x0e, 0x70, 0x65, 0x65, 0x72, 0x5f, 0x67, 0x72, 0x6f, 0x75, + 0x70, 0x5f, 0x69, 0x64, 0x73, 0x18, 0x05, 0x20, 0x03, 0x28, 0x09, 0x52, 0x0c, 0x70, 0x65, 0x65, + 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x73, 0x12, 0x28, 0x0a, 0x10, 0x70, 0x65, 0x65, + 0x72, 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x6e, 0x61, 0x6d, 0x65, 0x73, 0x18, 0x06, 0x20, + 0x03, 0x28, 0x09, 0x52, 0x0e, 0x70, 0x65, 0x65, 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x4e, 0x61, + 0x6d, 0x65, 0x73, 0x12, 0x23, 0x0a, 0x0d, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x5f, 0x74, + 0x6f, 0x6b, 0x65, 0x6e, 0x18, 0x07, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0c, 0x73, 0x65, 0x73, 0x73, + 0x69, 0x6f, 0x6e, 0x54, 0x6f, 0x6b, 0x65, 0x6e, 0x22, 0x50, 0x0a, 0x19, 0x56, 0x61, 0x6c, 0x69, + 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, + 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x1b, 0x0a, 0x09, 0x74, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x5f, + 0x69, 0x70, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x74, 0x75, 0x6e, 0x6e, 0x65, 0x6c, + 0x49, 0x70, 0x12, 0x16, 0x0a, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, 0x18, 0x02, 0x20, 0x01, + 0x28, 0x09, 0x52, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, 0x22, 0x84, 0x02, 0x0a, 0x1a, 0x56, + 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, 0x65, + 0x72, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x14, 0x0a, 0x05, 0x76, 0x61, 0x6c, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x08, 0x52, 0x05, 0x76, 0x61, 0x6c, 0x69, 0x64, 0x12, 0x17, 0x0a, 0x07, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, 0x75, 0x73, 0x65, 0x72, 0x49, 0x64, 0x12, 0x1d, 0x0a, 0x0a, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x65, 0x6d, 0x61, 0x69, 0x6c, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, 0x52, 0x09, 0x75, 0x73, 0x65, 0x72, 0x45, 0x6d, 0x61, 0x69, 0x6c, 0x12, 0x23, 0x0a, 0x0d, 0x64, 0x65, 0x6e, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x18, 0x04, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0c, - 0x64, 0x65, 0x6e, 0x69, 0x65, 0x64, 0x52, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x12, 0x24, 0x0a, 0x0e, - 0x70, 0x65, 0x65, 0x72, 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x73, 0x18, 0x05, - 0x20, 0x03, 0x28, 0x09, 0x52, 0x0c, 0x70, 0x65, 0x65, 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x49, - 0x64, 0x73, 0x12, 0x28, 0x0a, 0x10, 0x70, 0x65, 0x65, 0x72, 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, - 0x5f, 0x6e, 0x61, 0x6d, 0x65, 0x73, 0x18, 0x06, 0x20, 0x03, 0x28, 0x09, 0x52, 0x0e, 0x70, 0x65, - 0x65, 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x4e, 0x61, 0x6d, 0x65, 0x73, 0x22, 0x50, 0x0a, 0x19, - 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, - 0x65, 0x72, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x1b, 0x0a, 0x09, 0x74, 0x75, 0x6e, - 0x6e, 0x65, 0x6c, 0x5f, 0x69, 0x70, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x74, 0x75, - 0x6e, 0x6e, 0x65, 0x6c, 0x49, 0x70, 0x12, 0x16, 0x0a, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, - 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, 0x64, 0x6f, 0x6d, 0x61, 0x69, 0x6e, 0x22, 0x84, - 0x02, 0x0a, 0x1a, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, - 0x6c, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x14, 0x0a, - 0x05, 0x76, 0x61, 0x6c, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x08, 0x52, 0x05, 0x76, 0x61, - 0x6c, 0x69, 0x64, 0x12, 0x17, 0x0a, 0x07, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x69, 0x64, 0x18, 0x02, - 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, 0x75, 0x73, 0x65, 0x72, 0x49, 0x64, 0x12, 0x1d, 0x0a, 0x0a, - 0x75, 0x73, 0x65, 0x72, 0x5f, 0x65, 0x6d, 0x61, 0x69, 0x6c, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, - 0x52, 0x09, 0x75, 0x73, 0x65, 0x72, 0x45, 0x6d, 0x61, 0x69, 0x6c, 0x12, 0x23, 0x0a, 0x0d, 0x64, - 0x65, 0x6e, 0x69, 0x65, 0x64, 0x5f, 0x72, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x18, 0x04, 0x20, 0x01, - 0x28, 0x09, 0x52, 0x0c, 0x64, 0x65, 0x6e, 0x69, 0x65, 0x64, 0x52, 0x65, 0x61, 0x73, 0x6f, 0x6e, - 0x12, 0x23, 0x0a, 0x0d, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x5f, 0x74, 0x6f, 0x6b, 0x65, - 0x6e, 0x18, 0x05, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0c, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, - 0x54, 0x6f, 0x6b, 0x65, 0x6e, 0x12, 0x24, 0x0a, 0x0e, 0x70, 0x65, 0x65, 0x72, 0x5f, 0x67, 0x72, - 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x73, 0x18, 0x06, 0x20, 0x03, 0x28, 0x09, 0x52, 0x0c, 0x70, - 0x65, 0x65, 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x73, 0x12, 0x28, 0x0a, 0x10, 0x70, - 0x65, 0x65, 0x72, 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x6e, 0x61, 0x6d, 0x65, 0x73, 0x18, - 0x07, 0x20, 0x03, 0x28, 0x09, 0x52, 0x0e, 0x70, 0x65, 0x65, 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, - 0x4e, 0x61, 0x6d, 0x65, 0x73, 0x22, 0x81, 0x01, 0x0a, 0x13, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, - 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x32, 0x0a, - 0x04, 0x69, 0x6e, 0x69, 0x74, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1c, 0x2e, 0x6d, 0x61, - 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, - 0x70, 0x69, 0x6e, 0x67, 0x73, 0x49, 0x6e, 0x69, 0x74, 0x48, 0x00, 0x52, 0x04, 0x69, 0x6e, 0x69, - 0x74, 0x12, 0x2f, 0x0a, 0x03, 0x61, 0x63, 0x6b, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1b, - 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, - 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x41, 0x63, 0x6b, 0x48, 0x00, 0x52, 0x03, 0x61, - 0x63, 0x6b, 0x42, 0x05, 0x0a, 0x03, 0x6d, 0x73, 0x67, 0x22, 0xdf, 0x01, 0x0a, 0x10, 0x53, 0x79, - 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x49, 0x6e, 0x69, 0x74, 0x12, 0x19, - 0x0a, 0x08, 0x70, 0x72, 0x6f, 0x78, 0x79, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, - 0x52, 0x07, 0x70, 0x72, 0x6f, 0x78, 0x79, 0x49, 0x64, 0x12, 0x18, 0x0a, 0x07, 0x76, 0x65, 0x72, - 0x73, 0x69, 0x6f, 0x6e, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x07, 0x76, 0x65, 0x72, 0x73, - 0x69, 0x6f, 0x6e, 0x12, 0x39, 0x0a, 0x0a, 0x73, 0x74, 0x61, 0x72, 0x74, 0x65, 0x64, 0x5f, 0x61, - 0x74, 0x18, 0x03, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1a, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, - 0x2e, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x54, 0x69, 0x6d, 0x65, 0x73, 0x74, - 0x61, 0x6d, 0x70, 0x52, 0x09, 0x73, 0x74, 0x61, 0x72, 0x74, 0x65, 0x64, 0x41, 0x74, 0x12, 0x18, - 0x0a, 0x07, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x09, 0x52, - 0x07, 0x61, 0x64, 0x64, 0x72, 0x65, 0x73, 0x73, 0x12, 0x41, 0x0a, 0x0c, 0x63, 0x61, 0x70, 0x61, - 0x62, 0x69, 0x6c, 0x69, 0x74, 0x69, 0x65, 0x73, 0x18, 0x05, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1d, - 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x50, 0x72, 0x6f, 0x78, - 0x79, 0x43, 0x61, 0x70, 0x61, 0x62, 0x69, 0x6c, 0x69, 0x74, 0x69, 0x65, 0x73, 0x52, 0x0c, 0x63, - 0x61, 0x70, 0x61, 0x62, 0x69, 0x6c, 0x69, 0x74, 0x69, 0x65, 0x73, 0x22, 0x11, 0x0a, 0x0f, 0x53, - 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x41, 0x63, 0x6b, 0x22, 0x7e, - 0x0a, 0x14, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x52, 0x65, - 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x32, 0x0a, 0x07, 0x6d, 0x61, 0x70, 0x70, 0x69, 0x6e, - 0x67, 0x18, 0x01, 0x20, 0x03, 0x28, 0x0b, 0x32, 0x18, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, - 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, - 0x67, 0x52, 0x07, 0x6d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x12, 0x32, 0x0a, 0x15, 0x69, 0x6e, - 0x69, 0x74, 0x69, 0x61, 0x6c, 0x5f, 0x73, 0x79, 0x6e, 0x63, 0x5f, 0x63, 0x6f, 0x6d, 0x70, 0x6c, - 0x65, 0x74, 0x65, 0x18, 0x02, 0x20, 0x01, 0x28, 0x08, 0x52, 0x13, 0x69, 0x6e, 0x69, 0x74, 0x69, - 0x61, 0x6c, 0x53, 0x79, 0x6e, 0x63, 0x43, 0x6f, 0x6d, 0x70, 0x6c, 0x65, 0x74, 0x65, 0x22, 0xa9, - 0x01, 0x0a, 0x1b, 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, - 0x79, 0x4c, 0x69, 0x6d, 0x69, 0x74, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x1d, - 0x0a, 0x0a, 0x61, 0x63, 0x63, 0x6f, 0x75, 0x6e, 0x74, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, - 0x28, 0x09, 0x52, 0x09, 0x61, 0x63, 0x63, 0x6f, 0x75, 0x6e, 0x74, 0x49, 0x64, 0x12, 0x17, 0x0a, - 0x07, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, - 0x75, 0x73, 0x65, 0x72, 0x49, 0x64, 0x12, 0x1b, 0x0a, 0x09, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, - 0x69, 0x64, 0x73, 0x18, 0x03, 0x20, 0x03, 0x28, 0x09, 0x52, 0x08, 0x67, 0x72, 0x6f, 0x75, 0x70, - 0x49, 0x64, 0x73, 0x12, 0x1f, 0x0a, 0x0b, 0x70, 0x72, 0x6f, 0x76, 0x69, 0x64, 0x65, 0x72, 0x5f, - 0x69, 0x64, 0x18, 0x04, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0a, 0x70, 0x72, 0x6f, 0x76, 0x69, 0x64, - 0x65, 0x72, 0x49, 0x64, 0x12, 0x14, 0x0a, 0x05, 0x6d, 0x6f, 0x64, 0x65, 0x6c, 0x18, 0x05, 0x20, - 0x01, 0x28, 0x09, 0x52, 0x05, 0x6d, 0x6f, 0x64, 0x65, 0x6c, 0x22, 0xff, 0x01, 0x0a, 0x1c, 0x43, + 0x64, 0x65, 0x6e, 0x69, 0x65, 0x64, 0x52, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x12, 0x23, 0x0a, 0x0d, + 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x5f, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x18, 0x05, 0x20, + 0x01, 0x28, 0x09, 0x52, 0x0c, 0x73, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x54, 0x6f, 0x6b, 0x65, + 0x6e, 0x12, 0x24, 0x0a, 0x0e, 0x70, 0x65, 0x65, 0x72, 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, + 0x69, 0x64, 0x73, 0x18, 0x06, 0x20, 0x03, 0x28, 0x09, 0x52, 0x0c, 0x70, 0x65, 0x65, 0x72, 0x47, + 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x73, 0x12, 0x28, 0x0a, 0x10, 0x70, 0x65, 0x65, 0x72, 0x5f, + 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x6e, 0x61, 0x6d, 0x65, 0x73, 0x18, 0x07, 0x20, 0x03, 0x28, + 0x09, 0x52, 0x0e, 0x70, 0x65, 0x65, 0x72, 0x47, 0x72, 0x6f, 0x75, 0x70, 0x4e, 0x61, 0x6d, 0x65, + 0x73, 0x22, 0x81, 0x01, 0x0a, 0x13, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, + 0x67, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x32, 0x0a, 0x04, 0x69, 0x6e, 0x69, + 0x74, 0x18, 0x01, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1c, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, + 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, + 0x73, 0x49, 0x6e, 0x69, 0x74, 0x48, 0x00, 0x52, 0x04, 0x69, 0x6e, 0x69, 0x74, 0x12, 0x2f, 0x0a, + 0x03, 0x61, 0x63, 0x6b, 0x18, 0x02, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1b, 0x2e, 0x6d, 0x61, 0x6e, + 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, + 0x69, 0x6e, 0x67, 0x73, 0x41, 0x63, 0x6b, 0x48, 0x00, 0x52, 0x03, 0x61, 0x63, 0x6b, 0x42, 0x05, + 0x0a, 0x03, 0x6d, 0x73, 0x67, 0x22, 0xdf, 0x01, 0x0a, 0x10, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, + 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x49, 0x6e, 0x69, 0x74, 0x12, 0x19, 0x0a, 0x08, 0x70, 0x72, + 0x6f, 0x78, 0x79, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x07, 0x70, 0x72, + 0x6f, 0x78, 0x79, 0x49, 0x64, 0x12, 0x18, 0x0a, 0x07, 0x76, 0x65, 0x72, 0x73, 0x69, 0x6f, 0x6e, + 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x07, 0x76, 0x65, 0x72, 0x73, 0x69, 0x6f, 0x6e, 0x12, + 0x39, 0x0a, 0x0a, 0x73, 0x74, 0x61, 0x72, 0x74, 0x65, 0x64, 0x5f, 0x61, 0x74, 0x18, 0x03, 0x20, + 0x01, 0x28, 0x0b, 0x32, 0x1a, 0x2e, 0x67, 0x6f, 0x6f, 0x67, 0x6c, 0x65, 0x2e, 0x70, 0x72, 0x6f, + 0x74, 0x6f, 0x62, 0x75, 0x66, 0x2e, 0x54, 0x69, 0x6d, 0x65, 0x73, 0x74, 0x61, 0x6d, 0x70, 0x52, + 0x09, 0x73, 0x74, 0x61, 0x72, 0x74, 0x65, 0x64, 0x41, 0x74, 0x12, 0x18, 0x0a, 0x07, 0x61, 0x64, + 0x64, 0x72, 0x65, 0x73, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x09, 0x52, 0x07, 0x61, 0x64, 0x64, + 0x72, 0x65, 0x73, 0x73, 0x12, 0x41, 0x0a, 0x0c, 0x63, 0x61, 0x70, 0x61, 0x62, 0x69, 0x6c, 0x69, + 0x74, 0x69, 0x65, 0x73, 0x18, 0x05, 0x20, 0x01, 0x28, 0x0b, 0x32, 0x1d, 0x2e, 0x6d, 0x61, 0x6e, + 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x43, 0x61, 0x70, + 0x61, 0x62, 0x69, 0x6c, 0x69, 0x74, 0x69, 0x65, 0x73, 0x52, 0x0c, 0x63, 0x61, 0x70, 0x61, 0x62, + 0x69, 0x6c, 0x69, 0x74, 0x69, 0x65, 0x73, 0x22, 0x11, 0x0a, 0x0f, 0x53, 0x79, 0x6e, 0x63, 0x4d, + 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x41, 0x63, 0x6b, 0x22, 0x7e, 0x0a, 0x14, 0x53, 0x79, + 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, + 0x73, 0x65, 0x12, 0x32, 0x0a, 0x07, 0x6d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x18, 0x01, 0x20, + 0x03, 0x28, 0x0b, 0x32, 0x18, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, + 0x2e, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x52, 0x07, 0x6d, + 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x12, 0x32, 0x0a, 0x15, 0x69, 0x6e, 0x69, 0x74, 0x69, 0x61, + 0x6c, 0x5f, 0x73, 0x79, 0x6e, 0x63, 0x5f, 0x63, 0x6f, 0x6d, 0x70, 0x6c, 0x65, 0x74, 0x65, 0x18, + 0x02, 0x20, 0x01, 0x28, 0x08, 0x52, 0x13, 0x69, 0x6e, 0x69, 0x74, 0x69, 0x61, 0x6c, 0x53, 0x79, + 0x6e, 0x63, 0x43, 0x6f, 0x6d, 0x70, 0x6c, 0x65, 0x74, 0x65, 0x22, 0xa9, 0x01, 0x0a, 0x1b, 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x4c, 0x69, 0x6d, - 0x69, 0x74, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x1a, 0x0a, 0x08, 0x64, - 0x65, 0x63, 0x69, 0x73, 0x69, 0x6f, 0x6e, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x64, - 0x65, 0x63, 0x69, 0x73, 0x69, 0x6f, 0x6e, 0x12, 0x2c, 0x0a, 0x12, 0x73, 0x65, 0x6c, 0x65, 0x63, - 0x74, 0x65, 0x64, 0x5f, 0x70, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, - 0x01, 0x28, 0x09, 0x52, 0x10, 0x73, 0x65, 0x6c, 0x65, 0x63, 0x74, 0x65, 0x64, 0x50, 0x6f, 0x6c, - 0x69, 0x63, 0x79, 0x49, 0x64, 0x12, 0x30, 0x0a, 0x14, 0x61, 0x74, 0x74, 0x72, 0x69, 0x62, 0x75, - 0x74, 0x69, 0x6f, 0x6e, 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x18, 0x03, 0x20, - 0x01, 0x28, 0x09, 0x52, 0x12, 0x61, 0x74, 0x74, 0x72, 0x69, 0x62, 0x75, 0x74, 0x69, 0x6f, 0x6e, - 0x47, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x12, 0x25, 0x0a, 0x0e, 0x77, 0x69, 0x6e, 0x64, 0x6f, - 0x77, 0x5f, 0x73, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x03, 0x52, - 0x0d, 0x77, 0x69, 0x6e, 0x64, 0x6f, 0x77, 0x53, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x12, 0x1b, - 0x0a, 0x09, 0x64, 0x65, 0x6e, 0x79, 0x5f, 0x63, 0x6f, 0x64, 0x65, 0x18, 0x05, 0x20, 0x01, 0x28, - 0x09, 0x52, 0x08, 0x64, 0x65, 0x6e, 0x79, 0x43, 0x6f, 0x64, 0x65, 0x12, 0x1f, 0x0a, 0x0b, 0x64, - 0x65, 0x6e, 0x79, 0x5f, 0x72, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x18, 0x06, 0x20, 0x01, 0x28, 0x09, - 0x52, 0x0a, 0x64, 0x65, 0x6e, 0x79, 0x52, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x22, 0x91, 0x02, 0x0a, - 0x15, 0x52, 0x65, 0x63, 0x6f, 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, - 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x1d, 0x0a, 0x0a, 0x61, 0x63, 0x63, 0x6f, 0x75, 0x6e, - 0x74, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x09, 0x61, 0x63, 0x63, 0x6f, - 0x75, 0x6e, 0x74, 0x49, 0x64, 0x12, 0x17, 0x0a, 0x07, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x69, 0x64, - 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, 0x75, 0x73, 0x65, 0x72, 0x49, 0x64, 0x12, 0x19, - 0x0a, 0x08, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, - 0x52, 0x07, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x12, 0x25, 0x0a, 0x0e, 0x77, 0x69, 0x6e, - 0x64, 0x6f, 0x77, 0x5f, 0x73, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, - 0x03, 0x52, 0x0d, 0x77, 0x69, 0x6e, 0x64, 0x6f, 0x77, 0x53, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, - 0x12, 0x21, 0x0a, 0x0c, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x5f, 0x69, 0x6e, 0x70, 0x75, 0x74, - 0x18, 0x05, 0x20, 0x01, 0x28, 0x03, 0x52, 0x0b, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x49, 0x6e, - 0x70, 0x75, 0x74, 0x12, 0x23, 0x0a, 0x0d, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x5f, 0x6f, 0x75, - 0x74, 0x70, 0x75, 0x74, 0x18, 0x06, 0x20, 0x01, 0x28, 0x03, 0x52, 0x0c, 0x74, 0x6f, 0x6b, 0x65, - 0x6e, 0x73, 0x4f, 0x75, 0x74, 0x70, 0x75, 0x74, 0x12, 0x19, 0x0a, 0x08, 0x63, 0x6f, 0x73, 0x74, - 0x5f, 0x75, 0x73, 0x64, 0x18, 0x07, 0x20, 0x01, 0x28, 0x01, 0x52, 0x07, 0x63, 0x6f, 0x73, 0x74, - 0x55, 0x73, 0x64, 0x12, 0x1b, 0x0a, 0x09, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x73, - 0x18, 0x08, 0x20, 0x03, 0x28, 0x09, 0x52, 0x08, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x73, - 0x22, 0x18, 0x0a, 0x16, 0x52, 0x65, 0x63, 0x6f, 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, - 0x67, 0x65, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x2a, 0x64, 0x0a, 0x16, 0x50, 0x72, - 0x6f, 0x78, 0x79, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, - 0x54, 0x79, 0x70, 0x65, 0x12, 0x17, 0x0a, 0x13, 0x55, 0x50, 0x44, 0x41, 0x54, 0x45, 0x5f, 0x54, - 0x59, 0x50, 0x45, 0x5f, 0x43, 0x52, 0x45, 0x41, 0x54, 0x45, 0x44, 0x10, 0x00, 0x12, 0x18, 0x0a, - 0x14, 0x55, 0x50, 0x44, 0x41, 0x54, 0x45, 0x5f, 0x54, 0x59, 0x50, 0x45, 0x5f, 0x4d, 0x4f, 0x44, - 0x49, 0x46, 0x49, 0x45, 0x44, 0x10, 0x01, 0x12, 0x17, 0x0a, 0x13, 0x55, 0x50, 0x44, 0x41, 0x54, - 0x45, 0x5f, 0x54, 0x59, 0x50, 0x45, 0x5f, 0x52, 0x45, 0x4d, 0x4f, 0x56, 0x45, 0x44, 0x10, 0x02, - 0x2a, 0x46, 0x0a, 0x0f, 0x50, 0x61, 0x74, 0x68, 0x52, 0x65, 0x77, 0x72, 0x69, 0x74, 0x65, 0x4d, - 0x6f, 0x64, 0x65, 0x12, 0x18, 0x0a, 0x14, 0x50, 0x41, 0x54, 0x48, 0x5f, 0x52, 0x45, 0x57, 0x52, - 0x49, 0x54, 0x45, 0x5f, 0x44, 0x45, 0x46, 0x41, 0x55, 0x4c, 0x54, 0x10, 0x00, 0x12, 0x19, 0x0a, - 0x15, 0x50, 0x41, 0x54, 0x48, 0x5f, 0x52, 0x45, 0x57, 0x52, 0x49, 0x54, 0x45, 0x5f, 0x50, 0x52, - 0x45, 0x53, 0x45, 0x52, 0x56, 0x45, 0x10, 0x01, 0x2a, 0x90, 0x01, 0x0a, 0x0e, 0x4d, 0x69, 0x64, - 0x64, 0x6c, 0x65, 0x77, 0x61, 0x72, 0x65, 0x53, 0x6c, 0x6f, 0x74, 0x12, 0x1f, 0x0a, 0x1b, 0x4d, - 0x49, 0x44, 0x44, 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, 0x55, - 0x4e, 0x53, 0x50, 0x45, 0x43, 0x49, 0x46, 0x49, 0x45, 0x44, 0x10, 0x00, 0x12, 0x1e, 0x0a, 0x1a, - 0x4d, 0x49, 0x44, 0x44, 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, - 0x4f, 0x4e, 0x5f, 0x52, 0x45, 0x51, 0x55, 0x45, 0x53, 0x54, 0x10, 0x01, 0x12, 0x1f, 0x0a, 0x1b, - 0x4d, 0x49, 0x44, 0x44, 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, - 0x4f, 0x4e, 0x5f, 0x52, 0x45, 0x53, 0x50, 0x4f, 0x4e, 0x53, 0x45, 0x10, 0x02, 0x12, 0x1c, 0x0a, - 0x18, 0x4d, 0x49, 0x44, 0x44, 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, - 0x5f, 0x54, 0x45, 0x52, 0x4d, 0x49, 0x4e, 0x41, 0x4c, 0x10, 0x03, 0x2a, 0xc8, 0x01, 0x0a, 0x0b, - 0x50, 0x72, 0x6f, 0x78, 0x79, 0x53, 0x74, 0x61, 0x74, 0x75, 0x73, 0x12, 0x18, 0x0a, 0x14, 0x50, - 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x50, 0x45, 0x4e, 0x44, - 0x49, 0x4e, 0x47, 0x10, 0x00, 0x12, 0x17, 0x0a, 0x13, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, - 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x41, 0x43, 0x54, 0x49, 0x56, 0x45, 0x10, 0x01, 0x12, 0x23, - 0x0a, 0x1f, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x54, - 0x55, 0x4e, 0x4e, 0x45, 0x4c, 0x5f, 0x4e, 0x4f, 0x54, 0x5f, 0x43, 0x52, 0x45, 0x41, 0x54, 0x45, - 0x44, 0x10, 0x02, 0x12, 0x24, 0x0a, 0x20, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, - 0x54, 0x55, 0x53, 0x5f, 0x43, 0x45, 0x52, 0x54, 0x49, 0x46, 0x49, 0x43, 0x41, 0x54, 0x45, 0x5f, - 0x50, 0x45, 0x4e, 0x44, 0x49, 0x4e, 0x47, 0x10, 0x03, 0x12, 0x23, 0x0a, 0x1f, 0x50, 0x52, 0x4f, - 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x43, 0x45, 0x52, 0x54, 0x49, 0x46, - 0x49, 0x43, 0x41, 0x54, 0x45, 0x5f, 0x46, 0x41, 0x49, 0x4c, 0x45, 0x44, 0x10, 0x04, 0x12, 0x16, - 0x0a, 0x12, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x45, - 0x52, 0x52, 0x4f, 0x52, 0x10, 0x05, 0x32, 0xfc, 0x07, 0x0a, 0x0c, 0x50, 0x72, 0x6f, 0x78, 0x79, - 0x53, 0x65, 0x72, 0x76, 0x69, 0x63, 0x65, 0x12, 0x5f, 0x0a, 0x10, 0x47, 0x65, 0x74, 0x4d, 0x61, - 0x70, 0x70, 0x69, 0x6e, 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x12, 0x23, 0x2e, 0x6d, 0x61, - 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, 0x74, 0x4d, 0x61, 0x70, 0x70, - 0x69, 0x6e, 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, - 0x1a, 0x24, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, - 0x74, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, - 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x30, 0x01, 0x12, 0x55, 0x0a, 0x0c, 0x53, 0x79, 0x6e, 0x63, - 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x12, 0x1f, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, - 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, - 0x67, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x20, 0x2e, 0x6d, 0x61, 0x6e, 0x61, - 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, - 0x6e, 0x67, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x28, 0x01, 0x30, 0x01, 0x12, - 0x54, 0x0a, 0x0d, 0x53, 0x65, 0x6e, 0x64, 0x41, 0x63, 0x63, 0x65, 0x73, 0x73, 0x4c, 0x6f, 0x67, - 0x12, 0x20, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x65, - 0x6e, 0x64, 0x41, 0x63, 0x63, 0x65, 0x73, 0x73, 0x4c, 0x6f, 0x67, 0x52, 0x65, 0x71, 0x75, 0x65, - 0x73, 0x74, 0x1a, 0x21, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, - 0x53, 0x65, 0x6e, 0x64, 0x41, 0x63, 0x63, 0x65, 0x73, 0x73, 0x4c, 0x6f, 0x67, 0x52, 0x65, 0x73, - 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x51, 0x0a, 0x0c, 0x41, 0x75, 0x74, 0x68, 0x65, 0x6e, 0x74, - 0x69, 0x63, 0x61, 0x74, 0x65, 0x12, 0x1f, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, - 0x6e, 0x74, 0x2e, 0x41, 0x75, 0x74, 0x68, 0x65, 0x6e, 0x74, 0x69, 0x63, 0x61, 0x74, 0x65, 0x52, - 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x20, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, - 0x65, 0x6e, 0x74, 0x2e, 0x41, 0x75, 0x74, 0x68, 0x65, 0x6e, 0x74, 0x69, 0x63, 0x61, 0x74, 0x65, - 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x5d, 0x0a, 0x10, 0x53, 0x65, 0x6e, 0x64, - 0x53, 0x74, 0x61, 0x74, 0x75, 0x73, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x12, 0x23, 0x2e, 0x6d, - 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x65, 0x6e, 0x64, 0x53, 0x74, - 0x61, 0x74, 0x75, 0x73, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, - 0x74, 0x1a, 0x24, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, - 0x65, 0x6e, 0x64, 0x53, 0x74, 0x61, 0x74, 0x75, 0x73, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, - 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x5a, 0x0a, 0x0f, 0x43, 0x72, 0x65, 0x61, 0x74, - 0x65, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x50, 0x65, 0x65, 0x72, 0x12, 0x22, 0x2e, 0x6d, 0x61, 0x6e, + 0x69, 0x74, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x12, 0x1d, 0x0a, 0x0a, 0x61, 0x63, + 0x63, 0x6f, 0x75, 0x6e, 0x74, 0x5f, 0x69, 0x64, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x09, + 0x61, 0x63, 0x63, 0x6f, 0x75, 0x6e, 0x74, 0x49, 0x64, 0x12, 0x17, 0x0a, 0x07, 0x75, 0x73, 0x65, + 0x72, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, 0x06, 0x75, 0x73, 0x65, 0x72, + 0x49, 0x64, 0x12, 0x1b, 0x0a, 0x09, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x73, 0x18, + 0x03, 0x20, 0x03, 0x28, 0x09, 0x52, 0x08, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x73, 0x12, + 0x1f, 0x0a, 0x0b, 0x70, 0x72, 0x6f, 0x76, 0x69, 0x64, 0x65, 0x72, 0x5f, 0x69, 0x64, 0x18, 0x04, + 0x20, 0x01, 0x28, 0x09, 0x52, 0x0a, 0x70, 0x72, 0x6f, 0x76, 0x69, 0x64, 0x65, 0x72, 0x49, 0x64, + 0x12, 0x14, 0x0a, 0x05, 0x6d, 0x6f, 0x64, 0x65, 0x6c, 0x18, 0x05, 0x20, 0x01, 0x28, 0x09, 0x52, + 0x05, 0x6d, 0x6f, 0x64, 0x65, 0x6c, 0x22, 0xff, 0x01, 0x0a, 0x1c, 0x43, 0x68, 0x65, 0x63, 0x6b, + 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x4c, 0x69, 0x6d, 0x69, 0x74, 0x73, 0x52, + 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x1a, 0x0a, 0x08, 0x64, 0x65, 0x63, 0x69, 0x73, + 0x69, 0x6f, 0x6e, 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x64, 0x65, 0x63, 0x69, 0x73, + 0x69, 0x6f, 0x6e, 0x12, 0x2c, 0x0a, 0x12, 0x73, 0x65, 0x6c, 0x65, 0x63, 0x74, 0x65, 0x64, 0x5f, + 0x70, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, 0x01, 0x28, 0x09, 0x52, + 0x10, 0x73, 0x65, 0x6c, 0x65, 0x63, 0x74, 0x65, 0x64, 0x50, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x49, + 0x64, 0x12, 0x30, 0x0a, 0x14, 0x61, 0x74, 0x74, 0x72, 0x69, 0x62, 0x75, 0x74, 0x69, 0x6f, 0x6e, + 0x5f, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, 0x52, + 0x12, 0x61, 0x74, 0x74, 0x72, 0x69, 0x62, 0x75, 0x74, 0x69, 0x6f, 0x6e, 0x47, 0x72, 0x6f, 0x75, + 0x70, 0x49, 0x64, 0x12, 0x25, 0x0a, 0x0e, 0x77, 0x69, 0x6e, 0x64, 0x6f, 0x77, 0x5f, 0x73, 0x65, + 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x03, 0x52, 0x0d, 0x77, 0x69, 0x6e, + 0x64, 0x6f, 0x77, 0x53, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x12, 0x1b, 0x0a, 0x09, 0x64, 0x65, + 0x6e, 0x79, 0x5f, 0x63, 0x6f, 0x64, 0x65, 0x18, 0x05, 0x20, 0x01, 0x28, 0x09, 0x52, 0x08, 0x64, + 0x65, 0x6e, 0x79, 0x43, 0x6f, 0x64, 0x65, 0x12, 0x1f, 0x0a, 0x0b, 0x64, 0x65, 0x6e, 0x79, 0x5f, + 0x72, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x18, 0x06, 0x20, 0x01, 0x28, 0x09, 0x52, 0x0a, 0x64, 0x65, + 0x6e, 0x79, 0x52, 0x65, 0x61, 0x73, 0x6f, 0x6e, 0x22, 0x91, 0x02, 0x0a, 0x15, 0x52, 0x65, 0x63, + 0x6f, 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, + 0x73, 0x74, 0x12, 0x1d, 0x0a, 0x0a, 0x61, 0x63, 0x63, 0x6f, 0x75, 0x6e, 0x74, 0x5f, 0x69, 0x64, + 0x18, 0x01, 0x20, 0x01, 0x28, 0x09, 0x52, 0x09, 0x61, 0x63, 0x63, 0x6f, 0x75, 0x6e, 0x74, 0x49, + 0x64, 0x12, 0x17, 0x0a, 0x07, 0x75, 0x73, 0x65, 0x72, 0x5f, 0x69, 0x64, 0x18, 0x02, 0x20, 0x01, + 0x28, 0x09, 0x52, 0x06, 0x75, 0x73, 0x65, 0x72, 0x49, 0x64, 0x12, 0x19, 0x0a, 0x08, 0x67, 0x72, + 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x18, 0x03, 0x20, 0x01, 0x28, 0x09, 0x52, 0x07, 0x67, 0x72, + 0x6f, 0x75, 0x70, 0x49, 0x64, 0x12, 0x25, 0x0a, 0x0e, 0x77, 0x69, 0x6e, 0x64, 0x6f, 0x77, 0x5f, + 0x73, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x18, 0x04, 0x20, 0x01, 0x28, 0x03, 0x52, 0x0d, 0x77, + 0x69, 0x6e, 0x64, 0x6f, 0x77, 0x53, 0x65, 0x63, 0x6f, 0x6e, 0x64, 0x73, 0x12, 0x21, 0x0a, 0x0c, + 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x5f, 0x69, 0x6e, 0x70, 0x75, 0x74, 0x18, 0x05, 0x20, 0x01, + 0x28, 0x03, 0x52, 0x0b, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x49, 0x6e, 0x70, 0x75, 0x74, 0x12, + 0x23, 0x0a, 0x0d, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x5f, 0x6f, 0x75, 0x74, 0x70, 0x75, 0x74, + 0x18, 0x06, 0x20, 0x01, 0x28, 0x03, 0x52, 0x0c, 0x74, 0x6f, 0x6b, 0x65, 0x6e, 0x73, 0x4f, 0x75, + 0x74, 0x70, 0x75, 0x74, 0x12, 0x19, 0x0a, 0x08, 0x63, 0x6f, 0x73, 0x74, 0x5f, 0x75, 0x73, 0x64, + 0x18, 0x07, 0x20, 0x01, 0x28, 0x01, 0x52, 0x07, 0x63, 0x6f, 0x73, 0x74, 0x55, 0x73, 0x64, 0x12, + 0x1b, 0x0a, 0x09, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x5f, 0x69, 0x64, 0x73, 0x18, 0x08, 0x20, 0x03, + 0x28, 0x09, 0x52, 0x08, 0x67, 0x72, 0x6f, 0x75, 0x70, 0x49, 0x64, 0x73, 0x22, 0x18, 0x0a, 0x16, + 0x52, 0x65, 0x63, 0x6f, 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, + 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x2a, 0x64, 0x0a, 0x16, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x4d, + 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x54, 0x79, 0x70, 0x65, + 0x12, 0x17, 0x0a, 0x13, 0x55, 0x50, 0x44, 0x41, 0x54, 0x45, 0x5f, 0x54, 0x59, 0x50, 0x45, 0x5f, + 0x43, 0x52, 0x45, 0x41, 0x54, 0x45, 0x44, 0x10, 0x00, 0x12, 0x18, 0x0a, 0x14, 0x55, 0x50, 0x44, + 0x41, 0x54, 0x45, 0x5f, 0x54, 0x59, 0x50, 0x45, 0x5f, 0x4d, 0x4f, 0x44, 0x49, 0x46, 0x49, 0x45, + 0x44, 0x10, 0x01, 0x12, 0x17, 0x0a, 0x13, 0x55, 0x50, 0x44, 0x41, 0x54, 0x45, 0x5f, 0x54, 0x59, + 0x50, 0x45, 0x5f, 0x52, 0x45, 0x4d, 0x4f, 0x56, 0x45, 0x44, 0x10, 0x02, 0x2a, 0x46, 0x0a, 0x0f, + 0x50, 0x61, 0x74, 0x68, 0x52, 0x65, 0x77, 0x72, 0x69, 0x74, 0x65, 0x4d, 0x6f, 0x64, 0x65, 0x12, + 0x18, 0x0a, 0x14, 0x50, 0x41, 0x54, 0x48, 0x5f, 0x52, 0x45, 0x57, 0x52, 0x49, 0x54, 0x45, 0x5f, + 0x44, 0x45, 0x46, 0x41, 0x55, 0x4c, 0x54, 0x10, 0x00, 0x12, 0x19, 0x0a, 0x15, 0x50, 0x41, 0x54, + 0x48, 0x5f, 0x52, 0x45, 0x57, 0x52, 0x49, 0x54, 0x45, 0x5f, 0x50, 0x52, 0x45, 0x53, 0x45, 0x52, + 0x56, 0x45, 0x10, 0x01, 0x2a, 0x90, 0x01, 0x0a, 0x0e, 0x4d, 0x69, 0x64, 0x64, 0x6c, 0x65, 0x77, + 0x61, 0x72, 0x65, 0x53, 0x6c, 0x6f, 0x74, 0x12, 0x1f, 0x0a, 0x1b, 0x4d, 0x49, 0x44, 0x44, 0x4c, + 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, 0x55, 0x4e, 0x53, 0x50, 0x45, + 0x43, 0x49, 0x46, 0x49, 0x45, 0x44, 0x10, 0x00, 0x12, 0x1e, 0x0a, 0x1a, 0x4d, 0x49, 0x44, 0x44, + 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, 0x4f, 0x4e, 0x5f, 0x52, + 0x45, 0x51, 0x55, 0x45, 0x53, 0x54, 0x10, 0x01, 0x12, 0x1f, 0x0a, 0x1b, 0x4d, 0x49, 0x44, 0x44, + 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, 0x4f, 0x4e, 0x5f, 0x52, + 0x45, 0x53, 0x50, 0x4f, 0x4e, 0x53, 0x45, 0x10, 0x02, 0x12, 0x1c, 0x0a, 0x18, 0x4d, 0x49, 0x44, + 0x44, 0x4c, 0x45, 0x57, 0x41, 0x52, 0x45, 0x5f, 0x53, 0x4c, 0x4f, 0x54, 0x5f, 0x54, 0x45, 0x52, + 0x4d, 0x49, 0x4e, 0x41, 0x4c, 0x10, 0x03, 0x2a, 0xc8, 0x01, 0x0a, 0x0b, 0x50, 0x72, 0x6f, 0x78, + 0x79, 0x53, 0x74, 0x61, 0x74, 0x75, 0x73, 0x12, 0x18, 0x0a, 0x14, 0x50, 0x52, 0x4f, 0x58, 0x59, + 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x50, 0x45, 0x4e, 0x44, 0x49, 0x4e, 0x47, 0x10, + 0x00, 0x12, 0x17, 0x0a, 0x13, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, + 0x53, 0x5f, 0x41, 0x43, 0x54, 0x49, 0x56, 0x45, 0x10, 0x01, 0x12, 0x23, 0x0a, 0x1f, 0x50, 0x52, + 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x54, 0x55, 0x4e, 0x4e, 0x45, + 0x4c, 0x5f, 0x4e, 0x4f, 0x54, 0x5f, 0x43, 0x52, 0x45, 0x41, 0x54, 0x45, 0x44, 0x10, 0x02, 0x12, + 0x24, 0x0a, 0x20, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, + 0x43, 0x45, 0x52, 0x54, 0x49, 0x46, 0x49, 0x43, 0x41, 0x54, 0x45, 0x5f, 0x50, 0x45, 0x4e, 0x44, + 0x49, 0x4e, 0x47, 0x10, 0x03, 0x12, 0x23, 0x0a, 0x1f, 0x50, 0x52, 0x4f, 0x58, 0x59, 0x5f, 0x53, + 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x43, 0x45, 0x52, 0x54, 0x49, 0x46, 0x49, 0x43, 0x41, 0x54, + 0x45, 0x5f, 0x46, 0x41, 0x49, 0x4c, 0x45, 0x44, 0x10, 0x04, 0x12, 0x16, 0x0a, 0x12, 0x50, 0x52, + 0x4f, 0x58, 0x59, 0x5f, 0x53, 0x54, 0x41, 0x54, 0x55, 0x53, 0x5f, 0x45, 0x52, 0x52, 0x4f, 0x52, + 0x10, 0x05, 0x32, 0xfc, 0x07, 0x0a, 0x0c, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x53, 0x65, 0x72, 0x76, + 0x69, 0x63, 0x65, 0x12, 0x5f, 0x0a, 0x10, 0x47, 0x65, 0x74, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, + 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x12, 0x23, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, + 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, 0x74, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x55, + 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x24, 0x2e, 0x6d, + 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, 0x74, 0x4d, 0x61, 0x70, + 0x70, 0x69, 0x6e, 0x67, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, + 0x73, 0x65, 0x30, 0x01, 0x12, 0x55, 0x0a, 0x0c, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, + 0x69, 0x6e, 0x67, 0x73, 0x12, 0x1f, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, + 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x52, 0x65, + 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x20, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, + 0x6e, 0x74, 0x2e, 0x53, 0x79, 0x6e, 0x63, 0x4d, 0x61, 0x70, 0x70, 0x69, 0x6e, 0x67, 0x73, 0x52, + 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x28, 0x01, 0x30, 0x01, 0x12, 0x54, 0x0a, 0x0d, 0x53, + 0x65, 0x6e, 0x64, 0x41, 0x63, 0x63, 0x65, 0x73, 0x73, 0x4c, 0x6f, 0x67, 0x12, 0x20, 0x2e, 0x6d, + 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x65, 0x6e, 0x64, 0x41, 0x63, + 0x63, 0x65, 0x73, 0x73, 0x4c, 0x6f, 0x67, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x21, + 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x65, 0x6e, 0x64, + 0x41, 0x63, 0x63, 0x65, 0x73, 0x73, 0x4c, 0x6f, 0x67, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, + 0x65, 0x12, 0x51, 0x0a, 0x0c, 0x41, 0x75, 0x74, 0x68, 0x65, 0x6e, 0x74, 0x69, 0x63, 0x61, 0x74, + 0x65, 0x12, 0x1f, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x41, + 0x75, 0x74, 0x68, 0x65, 0x6e, 0x74, 0x69, 0x63, 0x61, 0x74, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, + 0x73, 0x74, 0x1a, 0x20, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, + 0x41, 0x75, 0x74, 0x68, 0x65, 0x6e, 0x74, 0x69, 0x63, 0x61, 0x74, 0x65, 0x52, 0x65, 0x73, 0x70, + 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x5d, 0x0a, 0x10, 0x53, 0x65, 0x6e, 0x64, 0x53, 0x74, 0x61, 0x74, + 0x75, 0x73, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x12, 0x23, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, + 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x65, 0x6e, 0x64, 0x53, 0x74, 0x61, 0x74, 0x75, 0x73, + 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x24, 0x2e, + 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x53, 0x65, 0x6e, 0x64, 0x53, + 0x74, 0x61, 0x74, 0x75, 0x73, 0x55, 0x70, 0x64, 0x61, 0x74, 0x65, 0x52, 0x65, 0x73, 0x70, 0x6f, + 0x6e, 0x73, 0x65, 0x12, 0x5a, 0x0a, 0x0f, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x50, 0x72, 0x6f, + 0x78, 0x79, 0x50, 0x65, 0x65, 0x72, 0x12, 0x22, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, + 0x65, 0x6e, 0x74, 0x2e, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x50, + 0x65, 0x65, 0x72, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x23, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x43, 0x72, 0x65, 0x61, 0x74, 0x65, 0x50, 0x72, - 0x6f, 0x78, 0x79, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x23, - 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x43, 0x72, 0x65, 0x61, - 0x74, 0x65, 0x50, 0x72, 0x6f, 0x78, 0x79, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, 0x73, 0x70, 0x6f, - 0x6e, 0x73, 0x65, 0x12, 0x4b, 0x0a, 0x0a, 0x47, 0x65, 0x74, 0x4f, 0x49, 0x44, 0x43, 0x55, 0x52, - 0x4c, 0x12, 0x1d, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, - 0x65, 0x74, 0x4f, 0x49, 0x44, 0x43, 0x55, 0x52, 0x4c, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, - 0x1a, 0x1e, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, - 0x74, 0x4f, 0x49, 0x44, 0x43, 0x55, 0x52, 0x4c, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, - 0x12, 0x5a, 0x0a, 0x0f, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, 0x73, - 0x69, 0x6f, 0x6e, 0x12, 0x22, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, + 0x6f, 0x78, 0x79, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, + 0x4b, 0x0a, 0x0a, 0x47, 0x65, 0x74, 0x4f, 0x49, 0x44, 0x43, 0x55, 0x52, 0x4c, 0x12, 0x1d, 0x2e, + 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, 0x74, 0x4f, 0x49, + 0x44, 0x43, 0x55, 0x52, 0x4c, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x1e, 0x2e, 0x6d, + 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x47, 0x65, 0x74, 0x4f, 0x49, 0x44, + 0x43, 0x55, 0x52, 0x4c, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x5a, 0x0a, 0x0f, + 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x12, + 0x22, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x56, 0x61, 0x6c, + 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x71, 0x75, + 0x65, 0x73, 0x74, 0x1a, 0x23, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, 0x73, 0x69, 0x6f, 0x6e, - 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x23, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, - 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x53, 0x65, 0x73, - 0x73, 0x69, 0x6f, 0x6e, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x63, 0x0a, 0x12, - 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, - 0x65, 0x72, 0x12, 0x25, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, - 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, - 0x65, 0x72, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x26, 0x2e, 0x6d, 0x61, 0x6e, 0x61, - 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, - 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, - 0x65, 0x12, 0x69, 0x0a, 0x14, 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, - 0x69, 0x63, 0x79, 0x4c, 0x69, 0x6d, 0x69, 0x74, 0x73, 0x12, 0x27, 0x2e, 0x6d, 0x61, 0x6e, 0x61, - 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, - 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x4c, 0x69, 0x6d, 0x69, 0x74, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, - 0x73, 0x74, 0x1a, 0x28, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, - 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x4c, 0x69, - 0x6d, 0x69, 0x74, 0x73, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x57, 0x0a, 0x0e, - 0x52, 0x65, 0x63, 0x6f, 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x12, 0x21, - 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x52, 0x65, 0x63, 0x6f, - 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, - 0x74, 0x1a, 0x22, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x52, - 0x65, 0x63, 0x6f, 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, 0x73, - 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x42, 0x08, 0x5a, 0x06, 0x2f, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x62, - 0x06, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x33, + 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x63, 0x0a, 0x12, 0x56, 0x61, 0x6c, 0x69, + 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, 0x65, 0x72, 0x12, 0x25, + 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x56, 0x61, 0x6c, 0x69, + 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, 0x6c, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, + 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x26, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, + 0x6e, 0x74, 0x2e, 0x56, 0x61, 0x6c, 0x69, 0x64, 0x61, 0x74, 0x65, 0x54, 0x75, 0x6e, 0x6e, 0x65, + 0x6c, 0x50, 0x65, 0x65, 0x72, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x69, 0x0a, + 0x14, 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x4c, + 0x69, 0x6d, 0x69, 0x74, 0x73, 0x12, 0x27, 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, + 0x6e, 0x74, 0x2e, 0x43, 0x68, 0x65, 0x63, 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, + 0x79, 0x4c, 0x69, 0x6d, 0x69, 0x74, 0x73, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x28, + 0x2e, 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x43, 0x68, 0x65, 0x63, + 0x6b, 0x4c, 0x4c, 0x4d, 0x50, 0x6f, 0x6c, 0x69, 0x63, 0x79, 0x4c, 0x69, 0x6d, 0x69, 0x74, 0x73, + 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x12, 0x57, 0x0a, 0x0e, 0x52, 0x65, 0x63, 0x6f, + 0x72, 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x12, 0x21, 0x2e, 0x6d, 0x61, 0x6e, + 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x52, 0x65, 0x63, 0x6f, 0x72, 0x64, 0x4c, 0x4c, + 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, 0x71, 0x75, 0x65, 0x73, 0x74, 0x1a, 0x22, 0x2e, + 0x6d, 0x61, 0x6e, 0x61, 0x67, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x2e, 0x52, 0x65, 0x63, 0x6f, 0x72, + 0x64, 0x4c, 0x4c, 0x4d, 0x55, 0x73, 0x61, 0x67, 0x65, 0x52, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, + 0x65, 0x42, 0x08, 0x5a, 0x06, 0x2f, 0x70, 0x72, 0x6f, 0x74, 0x6f, 0x62, 0x06, 0x70, 0x72, 0x6f, + 0x74, 0x6f, 0x33, } var ( diff --git a/shared/management/proto/proxy_service.proto b/shared/management/proto/proxy_service.proto index facadc4d5..7c5fb08eb 100644 --- a/shared/management/proto/proxy_service.proto +++ b/shared/management/proto/proxy_service.proto @@ -363,7 +363,9 @@ message GetOIDCURLResponse { message ValidateSessionRequest { string domain = 1; - string session_token = 2; + string session_token = 2 [deprecated = true]; + // session_code is a short-lived, single-use code exchanged for a session token. + string session_code = 3; } message ValidateSessionResponse { @@ -380,6 +382,8 @@ message ValidateSessionResponse { // Stamped onto upstream requests as X-NetBird-Groups so downstream // services can read names rather than opaque ids. repeated string peer_group_names = 6; + // session_token contains the durable token issued when session_code is redeemed. + string session_token = 7; } // ValidateTunnelPeerRequest carries the inbound peer's tunnel IP and the @@ -506,4 +510,3 @@ message RecordLLMUsageRequest { message RecordLLMUsageResponse { } - diff --git a/shared/management/types/networkmap_components.go b/shared/management/types/networkmap_components.go index e18db4ec0..7742ca244 100644 --- a/shared/management/types/networkmap_components.go +++ b/shared/management/types/networkmap_components.go @@ -58,6 +58,13 @@ type NetworkMapComponents struct { // domain targets. ForceRoutingPeerDNSResolution bool + // SkipRouteFirewallRules drops the route firewall rule computation from + // Calculate. A receiver without a firewall manager never reads + // RoutesFirewallRules, and on a routing peer with many network resources + // building them dominates the cost of a sync. Defaults to false so the + // management server keeps producing them. + SkipRouteFirewallRules bool + routesByPeerOnce sync.Once routesByPeerIdx map[string][]routeIndexEntry @@ -149,11 +156,15 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap { includeIPv6 = p.SupportsIPv6() && p.IPv6.IsValid() } routesUpdate := filterAndExpandRoutes(c.getRoutesToSync(targetPeerID, peersToConnect, peerGroups), includeIPv6) - routesFirewallRules := c.getPeerRoutesFirewallRules(ctx, targetPeerID, includeIPv6) + + var routesFirewallRules []*RouteFirewallRule + if !c.SkipRouteFirewallRules { + routesFirewallRules = c.getPeerRoutesFirewallRules(ctx, targetPeerID, includeIPv6) + } isRouter, networkResourcesRoutes, sourcePeers := c.getNetworkResourcesRoutesToSync(targetPeerID) var networkResourcesFirewallRules []*RouteFirewallRule - if isRouter { + if isRouter && !c.SkipRouteFirewallRules { networkResourcesFirewallRules = c.getPeerNetworkResourceFirewallRules(ctx, targetPeerID, networkResourcesRoutes, includeIPv6) } diff --git a/shared/profiling/profiling.go b/shared/profiling/profiling.go new file mode 100644 index 000000000..1d893048a --- /dev/null +++ b/shared/profiling/profiling.go @@ -0,0 +1,127 @@ +package profiling + +import ( + "errors" + "fmt" + "net/netip" + "net/url" + "os" + "strings" + "sync/atomic" + + "github.com/caarlos0/env/v11" + "github.com/grafana/pyroscope-go" + log "github.com/sirupsen/logrus" +) + +var errNotConfigured = errors.New("pyroscope not configured") + +var started atomic.Bool + +type config struct { + Address string `env:"NB_PYROSCOPE_ADDRESS"` + User string `env:"NB_PYROSCOPE_USER,notEmpty"` + Password string `env:"NB_PYROSCOPE_PASSWORD,notEmpty"` +} + +func Start(applicationName string) func() { + noop := func() {} + + cfg, err := loadConfig() + switch { + case errors.Is(err, errNotConfigured): + log.Info("pyroscope not configured, continuous profiling disabled") + return noop + case err != nil: + log.Errorf("failed to load pyroscope config: %v", err) + return noop + } + + // pprof allows one CPU profile per process, so a second profiler (e.g. the + // signal server inside the combined binary) would only log errors. + if !started.CompareAndSwap(false, true) { + log.Warnf("continuous profiling already running in this process, not starting it for %s", applicationName) + return noop + } + + tags := map[string]string{} + if hostname, err := os.Hostname(); err == nil { + tags["instance"] = hostname + } else { + log.Warnf("failed to resolve hostname for profile tags: %v", err) + } + + profiler, err := pyroscope.Start(pyroscope.Config{ + ApplicationName: applicationName, + ServerAddress: cfg.Address, + BasicAuthUser: cfg.User, + BasicAuthPassword: cfg.Password, + Logger: log.StandardLogger(), + Tags: tags, + ProfileTypes: []pyroscope.ProfileType{ + pyroscope.ProfileCPU, + pyroscope.ProfileAllocObjects, + pyroscope.ProfileAllocSpace, + pyroscope.ProfileInuseObjects, + pyroscope.ProfileInuseSpace, + }, + }) + if err != nil { + started.Store(false) + log.Errorf("failed to start continuous profiling: %v", err) + return noop + } + + return func() { + _ = profiler.Stop() + started.Store(false) + } +} + +func loadConfig() (config, error) { + var cfg config + if err := env.Parse(&cfg); err != nil { + if cfg.Address == "" { + return cfg, errNotConfigured + } + return cfg, fmt.Errorf("failed to parse pyroscope config: %w", err) + } + + if cfg.Address == "" { + return cfg, errNotConfigured + } + if err := validateAddress(cfg.Address); err != nil { + return cfg, err + } + + return cfg, nil +} + +// validateAddress refuses to send the basic-auth credentials in plaintext to +// anything but a loopback or private endpoint. +func validateAddress(address string) error { + u, err := url.Parse(address) + if err != nil { + return fmt.Errorf("invalid pyroscope address %q: %w", address, err) + } + + switch u.Scheme { + case "https": + return nil + case "http": + if isLocalOrPrivate(u.Hostname()) { + return nil + } + return fmt.Errorf("insecure pyroscope address %q: use https for non-local endpoints", address) + default: + return fmt.Errorf("pyroscope address %q must use http or https", address) + } +} + +func isLocalOrPrivate(host string) bool { + if host == "localhost" || strings.HasSuffix(host, ".localhost") { + return true + } + ip, err := netip.ParseAddr(host) + return err == nil && (ip.IsLoopback() || ip.IsPrivate()) +} diff --git a/shared/profiling/profiling_test.go b/shared/profiling/profiling_test.go new file mode 100644 index 000000000..68e56bb4c --- /dev/null +++ b/shared/profiling/profiling_test.go @@ -0,0 +1,202 @@ +package profiling + +import ( + "os" + "testing" + + log "github.com/sirupsen/logrus" + logtest "github.com/sirupsen/logrus/hooks/test" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestStartSkipsSecondProfilerInProcess(t *testing.T) { + clearEnv(t) + t.Setenv("NB_PYROSCOPE_ADDRESS", "http://127.0.0.1:1") + t.Setenv("NB_PYROSCOPE_USER", "user") + t.Setenv("NB_PYROSCOPE_PASSWORD", "token") + + started.Store(true) + t.Cleanup(func() { started.Store(false) }) + hook := logtest.NewGlobal() + t.Cleanup(hook.Reset) + + stop := Start("netbird-second") + stop() + + assert.True(t, started.Load(), "the running profiler must stay marked as started") + entry := hook.LastEntry() + require.NotNil(t, entry, "the skipped start must be logged") + assert.Equal(t, log.WarnLevel, entry.Level) + assert.Contains(t, entry.Message, "already running") +} + +func TestLoadConfig(t *testing.T) { + tests := []struct { + name string + env map[string]string + expected config + errIs error + wantErr bool + }{ + { + name: "address unset disables profiling", + errIs: errNotConfigured, + }, + { + name: "empty address disables profiling", + env: map[string]string{"NB_PYROSCOPE_ADDRESS": ""}, + errIs: errNotConfigured, + }, + { + name: "credentials without address disable profiling", + env: map[string]string{ + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + errIs: errNotConfigured, + }, + { + name: "address without credentials fails", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net", + }, + wantErr: true, + }, + { + name: "address with empty credentials fails", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net", + "NB_PYROSCOPE_USER": "", + "NB_PYROSCOPE_PASSWORD": "", + }, + wantErr: true, + }, + { + name: "address without password fails", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net", + "NB_PYROSCOPE_USER": "123456", + }, + wantErr: true, + }, + { + name: "full configuration", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "https://profiles-prod-001.grafana.net", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + expected: config{ + Address: "https://profiles-prod-001.grafana.net", + User: "123456", + Password: "token", + }, + }, + { + name: "http to loopback is allowed", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "http://127.0.0.1:4040", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + expected: config{ + Address: "http://127.0.0.1:4040", + User: "123456", + Password: "token", + }, + }, + { + name: "http to localhost is allowed", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "http://localhost:4040", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + expected: config{ + Address: "http://localhost:4040", + User: "123456", + Password: "token", + }, + }, + { + name: "http to private network is allowed", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "http://10.0.0.5:4040", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + expected: config{ + Address: "http://10.0.0.5:4040", + User: "123456", + Password: "token", + }, + }, + { + name: "http to public host is rejected", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "http://pyroscope.example.com", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + wantErr: true, + }, + { + name: "http to public address is rejected", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "http://203.0.113.10:4040", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + wantErr: true, + }, + { + name: "address without scheme is rejected", + env: map[string]string{ + "NB_PYROSCOPE_ADDRESS": "pyroscope.example.com:4040", + "NB_PYROSCOPE_USER": "123456", + "NB_PYROSCOPE_PASSWORD": "token", + }, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + clearEnv(t) + for k, v := range tt.env { + t.Setenv(k, v) + } + + cfg, err := loadConfig() + + switch { + case tt.errIs != nil: + require.ErrorIs(t, err, tt.errIs) + case tt.wantErr: + require.Error(t, err) + require.NotErrorIs(t, err, errNotConfigured) + default: + require.NoError(t, err) + assert.Equal(t, tt.expected, cfg) + } + }) + } +} + +func TestStartWithoutConfigurationIsNoop(t *testing.T) { + clearEnv(t) + + stop := Start("netbird-test") + require.NotNil(t, stop) + stop() +} + +func clearEnv(t *testing.T) { + t.Helper() + + for _, k := range []string{"NB_PYROSCOPE_ADDRESS", "NB_PYROSCOPE_USER", "NB_PYROSCOPE_PASSWORD"} { + t.Setenv(k, "") + require.NoError(t, os.Unsetenv(k)) + } +} diff --git a/management/server/http/middleware/rate_limiter.go b/shared/ratelimit/rate_limiter.go similarity index 77% rename from management/server/http/middleware/rate_limiter.go rename to shared/ratelimit/rate_limiter.go index bfd44afee..ff9a123a3 100644 --- a/management/server/http/middleware/rate_limiter.go +++ b/shared/ratelimit/rate_limiter.go @@ -1,7 +1,8 @@ -package middleware +package ratelimit import ( "context" + "encoding/json" "net" "net/http" "os" @@ -13,13 +14,14 @@ import ( log "github.com/sirupsen/logrus" "golang.org/x/time/rate" - "github.com/netbirdio/netbird/shared/management/http/util" + "github.com/netbirdio/netbird/trustedproxy" ) const ( - RateLimitingEnabledEnv = "NB_API_RATE_LIMITING_ENABLED" - RateLimitingBurstEnv = "NB_API_RATE_LIMITING_BURST" - RateLimitingRPMEnv = "NB_API_RATE_LIMITING_RPM" + RateLimitingEnabledEnv = "NB_API_RATE_LIMITING_ENABLED" + RateLimitingBurstEnv = "NB_API_RATE_LIMITING_BURST" + RateLimitingRPMEnv = "NB_API_RATE_LIMITING_RPM" + RateLimitingTrustedProxiesEnv = "NB_API_RATE_LIMITING_TRUSTED_PROXIES" defaultAPIRPM = 6 defaultAPIBurst = 500 @@ -35,6 +37,9 @@ type RateLimiterConfig struct { CleanupInterval time.Duration // LimiterTTL defines how long a limiter should be kept after last use (age threshold for removal) LimiterTTL time.Duration + // TrustedProxies lists the upstream proxies whose forwarding headers may be + // believed. Empty means requests are keyed by their direct peer address. + TrustedProxies *trustedproxy.List } // DefaultRateLimiterConfig returns a default configuration @@ -76,11 +81,18 @@ func RateLimiterConfigFromEnv() (cfg *RateLimiterConfig, enabled bool) { burst = defaultAPIBurst } + trusted, err := trustedproxy.Parse(os.Getenv(RateLimitingTrustedProxiesEnv)) + if err != nil { + log.Warnf("parsing %s env var: %v, trusting no proxies", RateLimitingTrustedProxiesEnv, err) + trusted = nil + } + return &RateLimiterConfig{ RequestsPerMinute: float64(rpm), Burst: burst, CleanupInterval: 6 * time.Hour, LimiterTTL: 24 * time.Hour, + TrustedProxies: trusted, }, os.Getenv(RateLimitingEnabledEnv) == "true" } @@ -250,17 +262,42 @@ func (rl *APIRateLimiter) Middleware(next http.Handler) http.Handler { next.ServeHTTP(w, r) return } - clientIP := getClientIP(r) + clientIP := getClientIP(r, rl.config.TrustedProxies) if !rl.Allow(clientIP) { - util.WriteErrorResponse("rate limit exceeded, please try again later", http.StatusTooManyRequests, w) + writeTooManyRequests(w) return } next.ServeHTTP(w, r) }) } -// getClientIP extracts the client IP address from the request. -func getClientIP(r *http.Request) string { +// errorResponse is the JSON body of a rejected request. +type errorResponse struct { + Message string `json:"message"` + Code int `json:"code"` +} + +// writeTooManyRequests writes a JSON error response with status 429 Too Many Requests. +func writeTooManyRequests(w http.ResponseWriter) { + w.Header().Set("Content-Type", "application/json; charset=UTF-8") + w.WriteHeader(http.StatusTooManyRequests) + if err := json.NewEncoder(w).Encode(errorResponse{ + Message: "rate limit exceeded, please try again later", + Code: http.StatusTooManyRequests, + }); err != nil { + log.Debugf("writing rate limit response: %v", err) + } +} + +// getClientIP extracts the client IP address from the request. Forwarding headers +// are used only when the request arrives from a trusted proxy. +func getClientIP(r *http.Request, trusted *trustedproxy.List) string { + if !trusted.Empty() { + if addr := trusted.ResolveClientIP(r.RemoteAddr, r.Header.Get("X-Forwarded-For")); addr.IsValid() { + return addr.String() + } + } + ip, _, err := net.SplitHostPort(r.RemoteAddr) if err != nil { return r.RemoteAddr diff --git a/management/server/http/middleware/rate_limiter_test.go b/shared/ratelimit/rate_limiter_test.go similarity index 85% rename from management/server/http/middleware/rate_limiter_test.go rename to shared/ratelimit/rate_limiter_test.go index 4b97d1874..c43b8e03d 100644 --- a/management/server/http/middleware/rate_limiter_test.go +++ b/shared/ratelimit/rate_limiter_test.go @@ -1,4 +1,4 @@ -package middleware +package ratelimit import ( "fmt" @@ -9,6 +9,9 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/trustedproxy" ) func TestAPIRateLimiter_Allow(t *testing.T) { @@ -63,6 +66,8 @@ func TestAPIRateLimiter_Middleware(t *testing.T) { rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) assert.Equal(t, http.StatusTooManyRequests, rr.Code) + assert.Equal(t, "application/json; charset=UTF-8", rr.Header().Get("Content-Type")) + assert.JSONEq(t, `{"message":"rate limit exceeded, please try again later","code":429}`, rr.Body.String()) } func TestAPIRateLimiter_Middleware_DifferentIPs(t *testing.T) { @@ -134,7 +139,7 @@ func TestGetClientIP(t *testing.T) { t.Run(tc.name, func(t *testing.T) { req := httptest.NewRequest(http.MethodGet, "/test", nil) req.RemoteAddr = tc.remoteAddr - assert.Equal(t, tc.expected, getClientIP(req)) + assert.Equal(t, tc.expected, getClientIP(req, nil)) }) } } @@ -327,3 +332,47 @@ func TestRateLimiterConfigFromEnv(t *testing.T) { assert.Equal(t, float64(defaultAPIRPM), cfg.RequestsPerMinute, "non-positive rpm must fall back to default") assert.Equal(t, defaultAPIBurst, cfg.Burst, "non-positive burst must fall back to default") } + +func TestGetClientIP_TrustedProxies(t *testing.T) { + trusted, err := trustedproxy.Parse("10.0.0.0/8") + require.NoError(t, err) + + tests := []struct { + name string + list *trustedproxy.List + remoteAddr string + xff string + expected string + }{ + { + name: "no trusted proxies ignores the header", + remoteAddr: "10.0.0.1:5555", + xff: "1.1.1.1, 2.2.2.2", + expected: "10.0.0.1", + }, + { + name: "behind a trusted proxy uses the right-most untrusted hop", + list: trusted, + remoteAddr: "10.0.0.1:5555", + xff: "1.1.1.1, 2.2.2.2", + expected: "2.2.2.2", + }, + { + name: "a caller reaching us directly cannot forge the header", + list: trusted, + remoteAddr: "203.0.113.5:5555", + xff: "1.1.1.1", + expected: "203.0.113.5", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/test", nil) + req.RemoteAddr = tc.remoteAddr + req.Header.Set("X-Forwarded-For", tc.xff) + + assert.Equal(t, tc.expected, getClientIP(req, tc.list)) + }) + } +} diff --git a/shared/relay/client/client.go b/shared/relay/client/client.go index 7171b40ad..9bb061f6e 100644 --- a/shared/relay/client/client.go +++ b/shared/relay/client/client.go @@ -30,6 +30,12 @@ const ( var ( ErrConnAlreadyExists = fmt.Errorf("connection already exists") + // ErrServerDisconnected is the cancellation cause of a relayed Conn when the + // client lost the connection to the relay server. + ErrServerDisconnected = fmt.Errorf("relay server disconnected") + // ErrPeerDisconnected is the cancellation cause of a relayed Conn when the + // remote peer went offline. + ErrPeerDisconnected = fmt.Errorf("remote peer disconnected") ) type internalStopFlag struct { @@ -74,16 +80,17 @@ type connContainer struct { msgChanLock sync.Mutex closed bool // flag to check if channel is closed ctx context.Context - cancel context.CancelFunc + cancel context.CancelCauseFunc } func newConnContainer(log *log.Entry, c *Client, peerID messages.PeerID, instanceURL *RelayAddr) *connContainer { - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithCancelCause(context.Background()) msgChan := make(chan Msg, connChannelSize) cn := &Conn{ dstID: peerID, messageChan: msgChan, instanceURL: instanceURL, + ctx: ctx, } cc := &connContainer{ log: log, @@ -106,10 +113,6 @@ func newConnContainer(log *log.Entry, c *Client, peerID messages.PeerID, instanc return cc } -func (cc *connContainer) netConn() net.Conn { - return cc.conn -} - func (cc *connContainer) writeMsg(msg Msg) { cc.msgChanLock.Lock() defer cc.msgChanLock.Unlock() @@ -128,8 +131,8 @@ func (cc *connContainer) writeMsg(msg Msg) { } } -func (cc *connContainer) close() { - cc.cancel() +func (cc *connContainer) close(cause error) { + cc.cancel(cause) cc.msgChanLock.Lock() defer cc.msgChanLock.Unlock() @@ -293,12 +296,12 @@ func (c *Client) Connect(ctx context.Context) error { return nil } -// OpenConn create a new net.Conn for the destination peer ID. In case if the connection is in progress +// OpenConn create a new Conn for the destination peer ID. In case if the connection is in progress // to the relay server, the function will block until the connection is established or timed out. Otherwise, // it will return immediately. // It block until the server confirm the peer is online. // todo: what should happen if call with the same peerID with multiple times? -func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (net.Conn, error) { +func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (*Conn, error) { peerID := messages.HashID(dstPeerID) c.mu.Lock() @@ -335,7 +338,7 @@ func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (net.Conn, erro delete(c.conns, peerID) } c.mu.Unlock() - container.close() + container.close(err) return nil, err } @@ -345,13 +348,13 @@ func (c *Client) OpenConn(ctx context.Context, dstPeerID string) (net.Conn, erro delete(c.conns, peerID) } c.mu.Unlock() - container.close() + container.close(ErrServerDisconnected) return nil, fmt.Errorf("relay connection is not established") } c.mu.Unlock() c.log.Infof("remote peer is available: %s", peerID) - return container.netConn(), nil + return container.conn, nil } // ServerInstanceURL returns the address of the relay server. It could change after the close and reopen the connection. @@ -773,7 +776,7 @@ func (c *Client) serverInstanceAddress() (string, netip.Addr, error) { func (c *Client) closeAllConns() { for _, container := range c.conns { - container.close() + container.close(ErrServerDisconnected) } c.conns = make(map[messages.PeerID]*connContainer) @@ -793,7 +796,7 @@ func (c *Client) closeConnsByPeerID(peerIDs []messages.PeerID) { } container.log.Infof("remote peer has been disconnected, free up connection: %s", peerID) - container.close() + container.close(ErrPeerDisconnected) delete(c.conns, peerID) } @@ -821,7 +824,7 @@ func (c *Client) closeConn(containerRef *connContainer, id messages.PeerID) erro c.log.Infof("free up connection to peer: %s", id) delete(c.conns, id) - current.close() + current.close(net.ErrClosed) return nil } diff --git a/shared/relay/client/conn.go b/shared/relay/client/conn.go index 9e2279790..67767a2b9 100644 --- a/shared/relay/client/conn.go +++ b/shared/relay/client/conn.go @@ -1,6 +1,7 @@ package client import ( + "context" "net" "time" @@ -12,11 +13,20 @@ type Conn struct { dstID messages.PeerID messageChan chan Msg instanceURL *RelayAddr + ctx context.Context writeFn func(messages.PeerID, []byte) (int, error) closeFn func(messages.PeerID) error localAddrFn func() net.Addr } +// Context returns a context that is cancelled when the connection is torn down, +// either by Close or by the relay client losing the server connection. The +// cancellation cause carries the reason, see ErrServerDisconnected and +// ErrPeerDisconnected. +func (c *Conn) Context() context.Context { + return c.ctx +} + func (c *Conn) Write(p []byte) (n int, err error) { return c.writeFn(c.dstID, p) } diff --git a/shared/relay/client/manager.go b/shared/relay/client/manager.go index 367c6dfc5..fc69e8fea 100644 --- a/shared/relay/client/manager.go +++ b/shared/relay/client/manager.go @@ -1,12 +1,9 @@ package client import ( - "container/list" "context" "fmt" - "net" "net/netip" - "reflect" "sync" "time" @@ -43,8 +40,6 @@ func NewRelayTrack() *RelayTrack { } } -type OnServerCloseListener func() - // ManagerOption configures a Manager at construction time. type ManagerOption func(*Manager) @@ -91,7 +86,6 @@ type Manager struct { relayClients map[string]*RelayTrack relayClientsMutex sync.RWMutex - onDisconnectedListeners map[string]*list.List onReconnectedListenerFn func() listenerLock sync.Mutex @@ -126,10 +120,9 @@ func NewManager(ctx context.Context, serverURLs []string, peerID string, mtu uin ConnectionTimeout: defaultConnectionTimeout, TransportFallback: tf, }, - relayClients: make(map[string]*RelayTrack), - onDisconnectedListeners: make(map[string]*list.List), - cleanupInterval: relayCleanupInterval, - keepUnusedServerTime: keepUnusedServerTime, + relayClients: make(map[string]*RelayTrack), + cleanupInterval: relayCleanupInterval, + keepUnusedServerTime: keepUnusedServerTime, } for _, opt := range opts { opt(m) @@ -168,11 +161,11 @@ func (m *Manager) Serve() error { // OpenConn opens a connection to the given peer key. If the peer is on the same relay server, the connection will be // established via the relay server. If the peer is on a different relay server, the manager will establish a new -// connection to the relay server. It returns back with a net.Conn what represent the remote peer connection. +// connection to the relay server. It returns the relayed connection to the remote peer. // // serverIP, when valid and serverAddress is foreign, is used as a dial target if the FQDN-based dial fails. // Ignored for the local home-server path. TLS verification still uses the FQDN via SNI. -func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) { +func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (*Conn, error) { m.relayClientMu.RLock() defer m.relayClientMu.RUnlock() @@ -185,9 +178,7 @@ func (m *Manager) OpenConn(ctx context.Context, serverAddress, peerKey string, s return nil, err } - var ( - netConn net.Conn - ) + var netConn *Conn if !foreign { log.Debugf("open peer connection via permanent server: %s", peerKey) netConn, err = m.relayClient.OpenConn(ctx, peerKey) @@ -220,31 +211,6 @@ func (m *Manager) SetOnReconnectedListener(f func()) { m.onReconnectedListenerFn = f } -// AddCloseListener adds a listener to the given server instance address. The listener will be called if the connection -// closed. -func (m *Manager) AddCloseListener(serverAddress string, onClosedListener OnServerCloseListener) error { - m.relayClientMu.RLock() - defer m.relayClientMu.RUnlock() - - if m.relayClient == nil { - return ErrRelayClientNotConnected - } - - foreign, err := m.isForeignServer(serverAddress) - if err != nil { - return err - } - - var listenerAddr string - if foreign { - listenerAddr = serverAddress - } else { - listenerAddr = m.relayClient.connectionURL - } - m.addListener(listenerAddr, onClosedListener) - return nil -} - // RelayInstanceAddress returns the address and resolved IP of the permanent relay server. It could change if the // network connection is lost. The address is sent to the target peer to choose the common relay server for the // communication; the IP is sent alongside so remote peers can dial directly without their own DNS lookup. Both @@ -330,7 +296,7 @@ func (m *Manager) UpdateToken(token *relayAuth.Token) error { return m.tokenStore.UpdateToken(token) } -func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (net.Conn, error) { +func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string, serverIP netip.Addr) (*Conn, error) { // check if already has a connection to the desired relay server m.relayClientsMutex.RLock() rt, ok := m.relayClients[serverAddress] @@ -383,7 +349,7 @@ func (m *Manager) openConnVia(ctx context.Context, serverAddress, peerKey string // waiting for the dial started by another openConnVia call to finish. It waits // on rt.ready rather than the track lock, so it neither holds nor contends the // track lock across the dial. -func (m *Manager) openConnOnTrack(ctx context.Context, rt *RelayTrack, peerKey string) (net.Conn, error) { +func (m *Manager) openConnOnTrack(ctx context.Context, rt *RelayTrack, peerKey string) (*Conn, error) { select { case <-rt.ready: case <-ctx.Done(): @@ -428,8 +394,6 @@ func (m *Manager) onServerDisconnected(serverAddress string) { if !isHome { m.evictForeignRelay(serverAddress) } - - m.notifyOnDisconnectListeners(serverAddress) } func (m *Manager) evictForeignRelay(serverAddress string) { @@ -523,36 +487,6 @@ func (m *Manager) cleanUpUnusedRelays() { } } -func (m *Manager) addListener(serverAddress string, onClosedListener OnServerCloseListener) { - m.listenerLock.Lock() - defer m.listenerLock.Unlock() - l, ok := m.onDisconnectedListeners[serverAddress] - if !ok { - l = list.New() - } - for e := l.Front(); e != nil; e = e.Next() { - if reflect.ValueOf(e.Value).Pointer() == reflect.ValueOf(onClosedListener).Pointer() { - return - } - } - l.PushBack(onClosedListener) - m.onDisconnectedListeners[serverAddress] = l -} - -func (m *Manager) notifyOnDisconnectListeners(serverAddress string) { - m.listenerLock.Lock() - defer m.listenerLock.Unlock() - - l, ok := m.onDisconnectedListeners[serverAddress] - if !ok { - return - } - for e := l.Front(); e != nil; e = e.Next() { - go e.Value.(OnServerCloseListener)() - } - delete(m.onDisconnectedListeners, serverAddress) -} - func relayConnState(c *Client) RelayConnState { addr, err := c.ServerInstanceURL() if err != nil { diff --git a/shared/relay/client/manager_test.go b/shared/relay/client/manager_test.go index 9e964f688..4a7840dd7 100644 --- a/shared/relay/client/manager_test.go +++ b/shared/relay/client/manager_test.go @@ -2,7 +2,9 @@ package client import ( "context" + "errors" "fmt" + "net" "net/netip" "testing" "time" @@ -291,35 +293,29 @@ func TestForeignAutoClose(t *testing.T) { t.Fatalf("failed to serve manager: %s", err) } - // Set up a disconnect listener to track when foreign server disconnects foreignServerURL := toURL(srvCfg2)[0] - disconnected := make(chan struct{}) - onDisconnect := func() { - select { - case disconnected <- struct{}{}: - default: - } - } t.Log("open connection to another peer") if _, err = mgr.OpenConn(ctx, foreignServerURL, "anotherpeer", netip.Addr{}); err == nil { t.Fatalf("should have failed to open connection to another peer") } - // Add the disconnect listener after the connection attempt - if err := mgr.AddCloseListener(foreignServerURL, onDisconnect); err != nil { - t.Logf("failed to add close listener (expected if connection failed): %s", err) - } - - // Wait for cleanup to happen timeout := relayCleanupInterval + keepUnusedServerTime + 2*time.Second t.Logf("waiting for relay cleanup: %s", timeout) - - select { - case <-disconnected: - t.Log("foreign relay connection cleaned up successfully") - case <-time.After(timeout): - t.Log("timeout waiting for cleanup - this might be expected if connection never established") + deadline := time.After(timeout) + for { + mgr.relayClientsMutex.RLock() + _, tracked := mgr.relayClients[foreignServerURL] + mgr.relayClientsMutex.RUnlock() + if !tracked { + t.Log("foreign relay connection cleaned up successfully") + break + } + select { + case <-deadline: + t.Fatal("foreign relay was not cleaned up") + case <-time.After(200 * time.Millisecond): + } } t.Logf("closing manager") @@ -413,23 +409,24 @@ func waitForReady(ctx context.Context, m *Manager, timeout time.Duration) error return fmt.Errorf("manager not ready within %s", timeout) } -func TestNotifierDoubleAdd(t *testing.T) { +func toURL(address server.ListenerConfig) []string { + return []string{"rel://" + address.Address} +} + +func TestConnContextCancelledOnServerDisconnect(t *testing.T) { ctx := context.Background() - listenerCfg1 := server.ListenerConfig{ - Address: "localhost:52501", - } - srv, err := server.NewServer(newManagerTestServerConfig(listenerCfg1.Address)) + srvCfg := server.ListenerConfig{Address: "localhost:52601"} + srv, err := server.NewServer(newManagerTestServerConfig(srvCfg.Address)) if err != nil { t.Fatalf("failed to create server: %s", err) } errChan := make(chan error, 1) go func() { - if err := srv.Listen(listenerCfg1); err != nil { + if err := srv.Listen(srvCfg); err != nil { errChan <- err } }() - defer func() { if err := srv.Shutdown(ctx); err != nil { t.Errorf("failed to close server: %s", err) @@ -440,46 +437,106 @@ func TestNotifierDoubleAdd(t *testing.T) { t.Fatalf("failed to start server: %s", err) } - log.Debugf("connect by alice") mCtx, cancel := context.WithCancel(ctx) defer cancel() - clientBob := NewManager(mCtx, toURL(listenerCfg1), "bob", iface.DefaultMTU) - if err = clientBob.Serve(); err != nil { + mgrBob := NewManager(mCtx, toURL(srvCfg), "bob", iface.DefaultMTU) + if err := mgrBob.Serve(); err != nil { + t.Fatalf("failed to serve bob manager: %s", err) + } + + mgr := NewManager(mCtx, toURL(srvCfg), "alice", iface.DefaultMTU) + if err := mgr.Serve(); err != nil { t.Fatalf("failed to serve manager: %s", err) } - clientAlice := NewManager(mCtx, toURL(listenerCfg1), "alice", iface.DefaultMTU) - if err = clientAlice.Serve(); err != nil { + ra, _, err := mgr.RelayInstanceAddress() + if err != nil { + t.Fatalf("failed to get relay address: %s", err) + } + + relayedConn, err := mgr.OpenConn(ctx, ra, "bob", netip.Addr{}) + if err != nil { + t.Fatalf("failed to open conn: %s", err) + } + + select { + case <-relayedConn.Context().Done(): + t.Fatal("conn context cancelled while the relay is still up") + default: + } + + _ = mgr.relayClient.relayConn.Close() + + select { + case <-relayedConn.Context().Done(): + case <-time.After(15 * time.Second): + t.Fatal("conn context was not cancelled after the relay connection dropped") + } + + if cause := context.Cause(relayedConn.Context()); !errors.Is(cause, ErrServerDisconnected) { + t.Errorf("unexpected cancellation cause: %v, want %v", cause, ErrServerDisconnected) + } +} + +func TestConnContextCauseOnLocalClose(t *testing.T) { + ctx := context.Background() + + srvCfg := server.ListenerConfig{Address: "localhost:52602"} + srv, err := server.NewServer(newManagerTestServerConfig(srvCfg.Address)) + if err != nil { + t.Fatalf("failed to create server: %s", err) + } + errChan := make(chan error, 1) + go func() { + if err := srv.Listen(srvCfg); err != nil { + errChan <- err + } + }() + defer func() { + if err := srv.Shutdown(ctx); err != nil { + t.Errorf("failed to close server: %s", err) + } + }() + + if err := waitForServerToStart(errChan); err != nil { + t.Fatalf("failed to start server: %s", err) + } + + mCtx, cancel := context.WithCancel(ctx) + defer cancel() + + mgrBob := NewManager(mCtx, toURL(srvCfg), "bob", iface.DefaultMTU) + if err := mgrBob.Serve(); err != nil { + t.Fatalf("failed to serve bob manager: %s", err) + } + + mgr := NewManager(mCtx, toURL(srvCfg), "alice", iface.DefaultMTU) + if err := mgr.Serve(); err != nil { t.Fatalf("failed to serve manager: %s", err) } - conn1, err := clientAlice.OpenConn(ctx, clientAlice.ServerURLs()[0], "bob", netip.Addr{}) + ra, _, err := mgr.RelayInstanceAddress() if err != nil { - t.Fatalf("failed to bind channel: %s", err) + t.Fatalf("failed to get relay address: %s", err) } - fnCloseListener := OnServerCloseListener(func() { - log.Infof("close listener") - }) - - err = clientAlice.AddCloseListener(clientAlice.ServerURLs()[0], fnCloseListener) + relayedConn, err := mgr.OpenConn(ctx, ra, "bob", netip.Addr{}) if err != nil { - t.Fatalf("failed to add close listener: %s", err) + t.Fatalf("failed to open conn: %s", err) } - err = clientAlice.AddCloseListener(clientAlice.ServerURLs()[0], fnCloseListener) - if err != nil { - t.Fatalf("failed to add close listener: %s", err) + if err := relayedConn.Close(); err != nil { + t.Fatalf("failed to close conn: %s", err) } - err = conn1.Close() - if err != nil { - t.Errorf("failed to close connection: %s", err) + select { + case <-relayedConn.Context().Done(): + case <-time.After(5 * time.Second): + t.Fatal("conn context was not cancelled after a local close") } -} - -func toURL(address server.ListenerConfig) []string { - return []string{"rel://" + address.Address} + if cause := context.Cause(relayedConn.Context()); !errors.Is(cause, net.ErrClosed) { + t.Errorf("unexpected cancellation cause after a local close: %v, want %v", cause, net.ErrClosed) + } } diff --git a/shared/signal/client/grpc.go b/shared/signal/client/grpc.go index a0bb2f080..92be57b30 100644 --- a/shared/signal/client/grpc.go +++ b/shared/signal/client/grpc.go @@ -53,6 +53,7 @@ type ConnStateNotifier interface { // GrpcClient Wraps the Signal Exchange Service gRpc client type GrpcClient struct { key wgtypes.Key + sharedKeys *encryption.SharedKeyCache realClient proto.SignalExchangeClient signalConn *grpc.ClientConn ctx context.Context @@ -107,6 +108,7 @@ func NewClient(ctx context.Context, addr string, key wgtypes.Key, tlsEnabled boo c := &GrpcClient{ ctx: ctx, key: key, + sharedKeys: encryption.NewSharedKeyCache(key), mux: sync.Mutex{}, status: StreamDisconnected, connStateCallbackLock: sync.RWMutex{}, @@ -158,6 +160,7 @@ func (c *GrpcClient) Close() error { } c.decryptionWg.Wait() c.decryptionWorker = nil + c.sharedKeys.Close() return c.signalConn.Close() } @@ -418,7 +421,7 @@ func (c *GrpcClient) decryptMessage(msg *proto.EncryptedMessage) (*proto.Message } body := &proto.Body{} - err = encryption.DecryptMessage(remoteKey, c.key, msg.GetBody(), body) + err = c.sharedKeys.DecryptMessage(remoteKey, msg.GetBody(), body) if err != nil { return nil, err } @@ -438,7 +441,7 @@ func (c *GrpcClient) encryptMessage(msg *proto.Message) (*proto.EncryptedMessage return nil, err } - encryptedBody, err := encryption.EncryptMessage(remoteKey, c.key, msg.Body) + encryptedBody, err := c.sharedKeys.EncryptMessage(remoteKey, msg.Body) if err != nil { return nil, err } diff --git a/sharedsock/sock_linux.go b/sharedsock/sock_linux.go index 150e8a722..640f813c9 100644 --- a/sharedsock/sock_linux.go +++ b/sharedsock/sock_linux.go @@ -10,13 +10,13 @@ import ( "context" "fmt" "net" + "net/netip" "time" "github.com/google/gopacket" "github.com/google/gopacket/layers" "github.com/mdlayher/socket" log "github.com/sirupsen/logrus" - "github.com/vishvananda/netlink" "golang.org/x/sync/errgroup" "golang.org/x/sys/unix" @@ -33,6 +33,8 @@ type SharedSocket struct { ctx context.Context conn4 *socket.Conn conn6 *socket.Conn + probe4 *srcProbe + probe6 *srcProbe port int mtu uint16 packetDemux chan rcvdPacket @@ -87,14 +89,26 @@ func Listen(port int, filter BPFFilter, mtu uint16) (_ net.PacketConn, err error return nil, fmt.Errorf("set SO_MARK on ipv4 socket: %w", err) } + if rawSock.probe4, err = newSrcProbe(unix.AF_INET); err != nil { + return nil, err + } + var sockErr error rawSock.conn6, sockErr = socket.Socket(unix.AF_INET6, unix.SOCK_RAW, unix.IPPROTO_UDP, "raw_udp6", nil) if sockErr != nil { - log.Errorf("Failed to create ipv6 raw socket: %v", err) + log.Errorf("Failed to create ipv6 raw socket: %v", sockErr) } else { if err = nbnet.SetSocketMark(rawSock.conn6); err != nil { return nil, fmt.Errorf("set SO_MARK on ipv6 socket: %w", err) } + rawSock.probe6, sockErr = newSrcProbe(unix.AF_INET6) + if sockErr != nil { + log.Errorf("Failed to create ipv6 source probe, continuing without ipv6: %v", sockErr) + if closeErr := rawSock.conn6.Close(); closeErr != nil { + log.Debugf("failed to close ipv6 raw socket: %v", closeErr) + } + rawSock.conn6 = nil + } } ipv4Instructions, ipv6Instructions, err := filter.GetInstructions(uint32(rawSock.port)) @@ -121,23 +135,36 @@ func Listen(port int, filter BPFFilter, mtu uint16) (_ net.PacketConn, err error return rawSock, nil } -// resolveSrc returns the source IP the kernel will pick for a packet sent to -// dst by these raw sockets, mirroring the fwmark the kernel will see on send. -func (s *SharedSocket) resolveSrc(dst net.IP) (net.IP, error) { - opts := &netlink.RouteGetOptions{} - if nbnet.AdvancedRouting() { - opts.Mark = nbnet.ControlPlaneMark +// sockaddr returns the raw send address for dst, carrying the scope of its zone. +func (s *SharedSocket) sockaddr(dst netip.Addr) (unix.Sockaddr, error) { + if dst.Zone() == "" { + return rawSockaddr(dst, 0), nil } - routes, err := netlink.RouteGetWithOptions(dst, opts) + if s.conn6 == nil { + return nil, fmt.Errorf("no raw socket for %s", dst) + } + rc, err := s.conn6.SyscallConn() if err != nil { - return nil, fmt.Errorf("route get %s: %w", dst, err) + return nil, fmt.Errorf("ipv6 raw socket: %w", err) } - for _, r := range routes { - if r.Src != nil { - return r.Src, nil - } + scope, err := zoneIndex(rc, dst.Zone()) + if err != nil { + return nil, err } - return nil, fmt.Errorf("no source IP for %s", dst) + return rawSockaddr(dst, scope), nil +} + +// resolveSrc returns the source IP the kernel will pick for a packet sent to sa +// by these raw sockets, mirroring the fwmark the kernel will see on send. +func (s *SharedSocket) resolveSrc(dst netip.Addr, sa unix.Sockaddr) (netip.Addr, error) { + probe := s.probe4 + if dst.Is6() { + probe = s.probe6 + } + if probe == nil { + return netip.Addr{}, fmt.Errorf("no raw socket for %s", dst) + } + return probe.resolve(sa) } // LocalAddr returns the local address, preferring IPv4 for backward compatibility. @@ -222,6 +249,13 @@ func (s *SharedSocket) Close() error { if s.conn6 != nil { errGrp.Go(s.conn6.Close) } + + if s.probe4 != nil { + errGrp.Go(s.probe4.close) + } + if s.probe6 != nil { + errGrp.Go(s.probe6.close) + } return errGrp.Wait() } @@ -296,14 +330,24 @@ func (s *SharedSocket) WriteTo(buf []byte, rAddr net.Addr) (n int, err error) { DstPort: layers.UDPPort(rUDPAddr.Port), } - src, err := s.resolveSrc(rUDPAddr.IP) - if err != nil { - return 0, fmt.Errorf("resolve source for %s: %w", rUDPAddr.IP, err) + dst := rUDPAddr.AddrPort().Addr().Unmap() + if !dst.IsValid() { + return 0, fmt.Errorf("invalid destination %s", rUDPAddr) } - rSockAddr, conn, nwLayer := s.getWriterObjects(src, rUDPAddr.IP) + rSockAddr, err := s.sockaddr(dst) + if err != nil { + return 0, err + } + + src, err := s.resolveSrc(dst, rSockAddr) + if err != nil { + return 0, fmt.Errorf("resolve source for %s: %w", dst, err) + } + + conn, nwLayer := s.getWriterObjects(src, dst) if conn == nil { - return 0, fmt.Errorf("no raw socket for %s", rUDPAddr.IP) + return 0, fmt.Errorf("no raw socket for %s", dst) } if err := udp.SetNetworkLayerForChecksum(nwLayer); err != nil { @@ -320,28 +364,23 @@ func (s *SharedSocket) WriteTo(buf []byte, rAddr net.Addr) (n int, err error) { } // getWriterObjects returns the specific IP version objects that are used to build a packet and send it using the raw socket -func (s *SharedSocket) getWriterObjects(src, dest net.IP) (sa unix.Sockaddr, conn *socket.Conn, layer gopacket.NetworkLayer) { - if dest.To4() == nil { - sa = &unix.SockaddrInet6{} - copy(sa.(*unix.SockaddrInet6).Addr[:], dest.To16()) +func (s *SharedSocket) getWriterObjects(src, dest netip.Addr) (conn *socket.Conn, layer gopacket.NetworkLayer) { + if dest.Is6() { conn = s.conn6 - layer = &layers.IPv6{ - SrcIP: src, - DstIP: dest, + SrcIP: src.AsSlice(), + DstIP: dest.AsSlice(), } } else { - sa = &unix.SockaddrInet4{} - copy(sa.(*unix.SockaddrInet4).Addr[:], dest.To4()) conn = s.conn4 layer = &layers.IPv4{ Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP, - SrcIP: src, - DstIP: dest, + SrcIP: src.AsSlice(), + DstIP: dest.AsSlice(), } } - return sa, conn, layer + return conn, layer } diff --git a/sharedsock/src_probe_linux.go b/sharedsock/src_probe_linux.go new file mode 100644 index 000000000..5e463b598 --- /dev/null +++ b/sharedsock/src_probe_linux.go @@ -0,0 +1,214 @@ +//go:build linux && !android + +package sharedsock + +import ( + "errors" + "fmt" + "net/netip" + "strconv" + "sync" + "syscall" + "unsafe" + + "github.com/mdlayher/socket" + log "github.com/sirupsen/logrus" + "golang.org/x/sys/unix" + + nbnet "github.com/netbirdio/netbird/client/net" +) + +var errProbeClosed = errors.New("source probe closed") + +// srcProbe finds the source address the kernel picks for a destination by connecting +// a UDP socket that carries the raw sockets' fwmark and reading back its local address. +// Connecting a UDP socket runs the output route lookup without sending anything. +type srcProbe struct { + family int + + mu sync.Mutex + // conn is nil while no socket is open. A failed route lookup keeps the socket, + // any other failure closes it and the next lookup opens a fresh one, so a socket + // in an unknown state is never reused. + conn *socket.Conn + closed bool +} + +// newSrcProbe opens a probe socket for the given address family. +func newSrcProbe(family int) (*srcProbe, error) { + conn, err := openProbeSocket(family) + if err != nil { + return nil, err + } + return &srcProbe{family: family, conn: conn}, nil +} + +// resolve returns the source address the kernel would use for a packet to sa, a +// sockaddr of the probe's family. It is safe for concurrent use. +func (p *srcProbe) resolve(sa unix.Sockaddr) (netip.Addr, error) { + p.mu.Lock() + defer p.mu.Unlock() + + if p.closed { + return netip.Addr{}, errProbeClosed + } + + if p.conn == nil { + conn, err := openProbeSocket(p.family) + if err != nil { + return netip.Addr{}, err + } + p.conn = conn + } + + src, err := p.lookup(sa) + if err != nil { + var rErr *routeError + if !errors.As(err, &rErr) { + if closeErr := p.closeSocket(); closeErr != nil { + log.Debugf("failed to close source probe socket: %v", closeErr) + } + } + return netip.Addr{}, err + } + return src, nil +} + +// close releases the socket. Later lookups fail with errProbeClosed. +func (p *srcProbe) close() error { + p.mu.Lock() + defer p.mu.Unlock() + + p.closed = true + return p.closeSocket() +} + +// closeSocket closes the socket if one is open. Callers must hold p.mu. +func (p *srcProbe) closeSocket() error { + conn := p.conn + p.conn = nil + if conn == nil { + return nil + } + return conn.Close() +} + +// lookup runs one route lookup on the socket. Callers must hold p.mu. +func (p *srcProbe) lookup(sa unix.Sockaddr) (netip.Addr, error) { + rc, err := p.conn.SyscallConn() + if err != nil { + return netip.Addr{}, fmt.Errorf("probe socket: %w", err) + } + + var src netip.Addr + var lookupErr error + if err := rc.Control(func(fd uintptr) { + src, lookupErr = lookupFD(int(fd), sa) + }); err != nil { + return netip.Addr{}, fmt.Errorf("probe socket: %w", err) + } + return src, lookupErr +} + +// routeError is a route lookup the kernel refused. The socket is still usable after it. +type routeError struct { + err error +} + +func (e *routeError) Error() string { + return fmt.Sprintf("route lookup: %v", e.err) +} + +func (e *routeError) Unwrap() error { + return e.err +} + +func lookupFD(fd int, sa unix.Sockaddr) (netip.Addr, error) { + // A connected socket keeps the source address of its first connect and reuses + // it for later route lookups, so dissolve the association first. + if err := disconnect(fd); err != nil { + return netip.Addr{}, fmt.Errorf("disconnect probe socket: %w", err) + } + + if err := unix.Connect(fd, sa); err != nil { + return netip.Addr{}, &routeError{err: err} + } + + local, err := unix.Getsockname(fd) + if err != nil { + return netip.Addr{}, fmt.Errorf("read probe socket address: %w", err) + } + + var src netip.Addr + switch a := local.(type) { + case *unix.SockaddrInet4: + src = netip.AddrFrom4(a.Addr) + case *unix.SockaddrInet6: + src = netip.AddrFrom16(a.Addr) + } + if !src.IsValid() || src.IsUnspecified() { + return netip.Addr{}, &routeError{err: errors.New("no source address")} + } + return src, nil +} + +func openProbeSocket(family int) (*socket.Conn, error) { + conn, err := socket.Socket(family, unix.SOCK_DGRAM, unix.IPPROTO_UDP, "udp_src_probe", nil) + if err != nil { + return nil, fmt.Errorf("create source probe socket: %w", err) + } + + if err := nbnet.SetSocketMark(conn); err != nil { + _ = conn.Close() + return nil, fmt.Errorf("set SO_MARK on source probe socket: %w", err) + } + return conn, nil +} + +// disconnect dissolves a UDP socket's association by connecting to AF_UNSPEC, which +// also clears the source address the kernel pinned on the previous connect. +func disconnect(fd int) error { + sa := unix.RawSockaddr{Family: unix.AF_UNSPEC} + _, _, errno := unix.Syscall(unix.SYS_CONNECT, uintptr(fd), uintptr(unsafe.Pointer(&sa)), unsafe.Sizeof(sa)) + if errno != 0 { + return errno + } + return nil +} + +// rawSockaddr returns the sockaddr for dst with port 0 and the given scope. Port 0 +// matches a raw send, whose route lookup carries no ports. A UDP probe connected +// to it still gets an ephemeral source port before its lookup. +func rawSockaddr(dst netip.Addr, scope uint32) unix.Sockaddr { + if dst.Is4() { + return &unix.SockaddrInet4{Addr: dst.As4()} + } + return &unix.SockaddrInet6{Addr: dst.As16(), ZoneId: scope} +} + +// zoneIndex returns the interface index for an IPv6 zone, which is either an +// interface name or a numeric index. An empty zone is index 0. A name costs one +// SIOCGIFINDEX ioctl on rc, which may be any socket. +func zoneIndex(rc syscall.RawConn, zone string) (uint32, error) { + if zone == "" { + return 0, nil + } + if idx, err := strconv.ParseUint(zone, 10, 32); err == nil { + return uint32(idx), nil + } + + ifr, err := unix.NewIfreq(zone) + if err != nil { + return 0, fmt.Errorf("zone %q: %w", zone, err) + } + var ioctlErr error + if err := rc.Control(func(fd uintptr) { + ioctlErr = unix.IoctlIfreq(int(fd), unix.SIOCGIFINDEX, ifr) + }); err != nil { + return 0, fmt.Errorf("zone %q: %w", zone, err) + } + if ioctlErr != nil { + return 0, fmt.Errorf("resolve zone %q: %w", zone, ioctlErr) + } + return ifr.Uint32(), nil +} diff --git a/sharedsock/src_probe_linux_test.go b/sharedsock/src_probe_linux_test.go new file mode 100644 index 000000000..bdf8e37d8 --- /dev/null +++ b/sharedsock/src_probe_linux_test.go @@ -0,0 +1,218 @@ +//go:build linux && !android + +package sharedsock + +import ( + "net" + "net/netip" + "strconv" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vishvananda/netlink" + "golang.org/x/sys/unix" +) + +// routeGetSrc is the kernel's answer through netlink, used as the reference. +func routeGetSrc(t *testing.T, dst netip.Addr) (netip.Addr, bool) { + t.Helper() + routes, err := netlink.RouteGet(net.IP(dst.AsSlice())) + if err != nil { + return netip.Addr{}, false + } + for _, r := range routes { + if src, ok := netip.AddrFromSlice(r.Src); ok { + return src.Unmap(), true + } + } + return netip.Addr{}, false +} + +func newTestProbe(t *testing.T, family int) *srcProbe { + t.Helper() + p, err := newSrcProbe(family) + require.NoError(t, err) + t.Cleanup(func() { _ = p.close() }) + return p +} + +// A reused UDP socket keeps the source address of its first connect. Alternating +// between a loopback and an off-host destination catches a probe that forgets to +// disconnect: the second lookup would report 127.0.0.1 for the off-host address. +func TestSrcProbe_AlternatingDestinationsMatchRouteGet(t *testing.T) { + loopback := netip.MustParseAddr("127.0.0.1") + remote := netip.MustParseAddr("192.0.2.1") + + remoteSrc, ok := routeGetSrc(t, remote) + if !ok { + t.Skip("no route to an off-host IPv4 destination") + } + require.NotEqual(t, loopback, remoteSrc, "off-host destination must not route via loopback") + + p := newTestProbe(t, unix.AF_INET) + for i := 0; i < 3; i++ { + src, err := p.resolve(rawSockaddr(loopback, 0)) + require.NoError(t, err) + assert.Equal(t, loopback, src, "source for loopback, round %d", i) + + src, err = p.resolve(rawSockaddr(remote, 0)) + require.NoError(t, err) + assert.Equal(t, remoteSrc, src, "source for %s, round %d", remote, i) + } +} + +func TestSrcProbe_IPv6(t *testing.T) { + loopback := netip.MustParseAddr("::1") + if _, ok := routeGetSrc(t, loopback); !ok { + t.Skip("no IPv6 loopback") + } + + p := newTestProbe(t, unix.AF_INET6) + src, err := p.resolve(rawSockaddr(loopback, 0)) + require.NoError(t, err) + assert.Equal(t, loopback, src, "source for ::1") + + remote := netip.MustParseAddr("2001:db8::1") + remoteSrc, ok := routeGetSrc(t, remote) + if !ok { + t.Skipf("no route to %s", remote) + } + src, err = p.resolve(rawSockaddr(remote, 0)) + require.NoError(t, err) + assert.Equal(t, remoteSrc, src, "source for %s", remote) +} + +// A route the kernel refuses is an answer, not a broken socket, so the probe keeps +// its socket and the next lookup reuses it. +func TestSrcProbe_RouteErrorKeepsSocket(t *testing.T) { + p := newTestProbe(t, unix.AF_INET) + conn := p.conn + + // Connecting to the limited broadcast address without SO_BROADCAST fails. + _, err := p.resolve(rawSockaddr(netip.MustParseAddr("255.255.255.255"), 0)) + var rErr *routeError + require.ErrorAs(t, err, &rErr, "lookup for the broadcast address should fail as a route error") + assert.Same(t, conn, p.conn, "route error should keep the probe socket") + + src, err := p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + require.NoError(t, err) + assert.Equal(t, netip.MustParseAddr("127.0.0.1"), src, "source after the route error") + assert.Same(t, conn, p.conn, "lookup after a route error should reuse the socket") +} + +// A socket that stops working is dropped, and the next lookup opens a fresh one. +func TestSrcProbe_ReopensAfterSocketError(t *testing.T) { + p := newTestProbe(t, unix.AF_INET) + require.NoError(t, p.conn.Close()) + + _, err := p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + require.Error(t, err, "lookup on a closed socket should fail") + assert.Nil(t, p.conn, "socket error should drop the probe socket") + + src, err := p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + require.NoError(t, err) + assert.Equal(t, netip.MustParseAddr("127.0.0.1"), src, "source after reopening") + assert.NotNil(t, p.conn, "probe socket should be open again") +} + +// A link-local destination is only routable with its scope. The probe must pass the +// scope to the kernel and get the interface's own link-local address back. +func TestSrcProbe_LinkLocalWithZone(t *testing.T) { + iface, want := linkLocalInterface(t) + dst := netip.MustParseAddr("fe80::1") + + p := newTestProbe(t, unix.AF_INET6) + rc, err := p.conn.SyscallConn() + require.NoError(t, err) + + for _, zone := range []string{iface.Name, strconv.Itoa(iface.Index)} { + scope, err := zoneIndex(rc, zone) + require.NoError(t, err, "zone %q", zone) + assert.Equal(t, uint32(iface.Index), scope, "index for zone %q", zone) + + src, err := p.resolve(rawSockaddr(dst.WithZone(zone), scope)) + require.NoError(t, err, "zone %q", zone) + assert.Equal(t, want, src, "source for %s%%%s", dst, zone) + } + + _, err = p.resolve(rawSockaddr(dst, 0)) + var rErr *routeError + assert.ErrorAs(t, err, &rErr, "link-local destination without a scope should fail as a route error") +} + +func TestZoneIndex(t *testing.T) { + p := newTestProbe(t, unix.AF_INET6) + rc, err := p.conn.SyscallConn() + require.NoError(t, err) + + scope, err := zoneIndex(rc, "") + require.NoError(t, err) + assert.Zero(t, scope, "empty zone should be index 0") + + _, err = zoneIndex(rc, "nb-no-such-if0") + assert.Error(t, err, "unknown interface name should fail") +} + +// linkLocalInterface returns an up interface and its IPv6 link-local address. +func linkLocalInterface(t *testing.T) (net.Interface, netip.Addr) { + t.Helper() + ifaces, err := net.Interfaces() + require.NoError(t, err) + for _, iface := range ifaces { + if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 { + continue + } + addrs, err := iface.Addrs() + if err != nil { + continue + } + for _, a := range addrs { + prefix, err := netip.ParsePrefix(a.String()) + if err == nil && prefix.Addr().IsLinkLocalUnicast() { + return iface, prefix.Addr() + } + } + } + t.Skip("no interface with an IPv6 link-local address") + return net.Interface{}, netip.Addr{} +} + +func TestSrcProbe_ClosedRejectsLookups(t *testing.T) { + p, err := newSrcProbe(unix.AF_INET) + require.NoError(t, err) + require.NoError(t, p.close()) + require.NoError(t, p.close(), "close must be idempotent") + + _, err = p.resolve(rawSockaddr(netip.MustParseAddr("127.0.0.1"), 0)) + assert.ErrorIs(t, err, errProbeClosed) + assert.Nil(t, p.conn, "closed probe must not reopen") +} + +func BenchmarkSrcProbe(b *testing.B) { + p, err := newSrcProbe(unix.AF_INET) + require.NoError(b, err) + defer p.close() + + dst := netip.MustParseAddr("192.0.2.1") + if _, err := p.resolve(rawSockaddr(dst, 0)); err != nil { + b.Skipf("no route to %s: %v", dst, err) + } + + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if _, err := p.resolve(rawSockaddr(dst, 0)); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkSrcRouteGet(b *testing.B) { + dst := net.ParseIP("192.0.2.1") + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if _, err := netlink.RouteGetWithOptions(dst, &netlink.RouteGetOptions{}); err != nil { + b.Skipf("no route to %s: %v", dst, err) + } + } +} diff --git a/sharedsock/src_probe_privileged_linux_test.go b/sharedsock/src_probe_privileged_linux_test.go new file mode 100644 index 000000000..fed1d3985 --- /dev/null +++ b/sharedsock/src_probe_privileged_linux_test.go @@ -0,0 +1,92 @@ +//go:build privileged + +package sharedsock + +import ( + "net" + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vishvananda/netlink" + "golang.org/x/sys/unix" + + nbnet "github.com/netbirdio/netbird/client/net" +) + +// The probe must pick the source the raw sockets get on send, which carry the +// control-plane fwmark. The test installs the same shape of policy rule the client +// uses: unmarked traffic to 192.0.2.0/24 is diverted into a table that routes it +// via loopback, so an unmarked lookup reports 127.0.0.1 and a marked one does not. +func TestSrcProbe_HonoursControlPlaneMark(t *testing.T) { + nbnet.Init() + if !nbnet.AdvancedRouting() { + t.Skip("advanced routing not supported") + } + + const ( + table = 4242 + // Below the client's own rules, so a default route in the netbird table cannot win. + priority = 90 + ) + dst := netip.MustParseAddr("192.0.2.1") + + marked, err := netlink.RouteGetWithOptions(net.IP(dst.AsSlice()), &netlink.RouteGetOptions{Mark: nbnet.ControlPlaneMark}) + if err != nil { + t.Skipf("no route to %s: %v", dst, err) + } + if len(marked) == 0 || marked[0].Src == nil { + t.Skipf("marked route to %s has no source address", dst) + } + markedSrc, ok := netip.AddrFromSlice(marked[0].Src) + require.True(t, ok, "parse marked source") + markedSrc = markedSrc.Unmap() + if markedSrc == netip.MustParseAddr("127.0.0.1") { + t.Skipf("marked route to %s already uses loopback, no contrast to test", dst) + } + + rules, err := netlink.RuleList(unix.AF_INET) + require.NoError(t, err) + for _, r := range rules { + if r.Priority == priority || r.Table == table { + t.Skipf("rule priority %d or table %d already in use", priority, table) + } + } + + lo, err := netlink.LinkByName("lo") + require.NoError(t, err) + + route := &netlink.Route{ + Dst: &net.IPNet{IP: net.IPv4(192, 0, 2, 0), Mask: net.CIDRMask(24, 32)}, + LinkIndex: lo.Attrs().Index, + // 127.0.0.1 is host-scoped, so a link-scoped route only picks it when told to. + Src: net.IPv4(127, 0, 0, 1), + Table: table, + Scope: netlink.SCOPE_LINK, + } + require.NoError(t, netlink.RouteAdd(route)) + t.Cleanup(func() { _ = netlink.RouteDel(route) }) + + rule := netlink.NewRule() + rule.Family = unix.AF_INET + rule.Priority = priority + rule.Table = table + rule.Mark = nbnet.ControlPlaneMark + rule.Invert = true + require.NoError(t, netlink.RuleAdd(rule)) + t.Cleanup(func() { _ = netlink.RuleDel(rule) }) + + unmarked, err := netlink.RouteGet(net.IP(dst.AsSlice())) + require.NoError(t, err) + require.NotEmpty(t, unmarked) + require.True(t, unmarked[0].Src.Equal(net.IPv4(127, 0, 0, 1)), "unmarked lookup should be diverted to loopback, got %s", unmarked[0].Src) + + p, err := newSrcProbe(unix.AF_INET) + require.NoError(t, err) + t.Cleanup(func() { _ = p.close() }) + + src, err := p.resolve(rawSockaddr(dst, 0)) + require.NoError(t, err) + assert.Equal(t, markedSrc, src, "probe should resolve the source of the marked lookup") +} diff --git a/signal/cmd/run.go b/signal/cmd/run.go index a36623c6b..42b7d2505 100644 --- a/signal/cmd/run.go +++ b/signal/cmd/run.go @@ -119,6 +119,7 @@ var ( if err != nil { return fmt.Errorf("creating signal server: %v", err) } + defer srv.Stop() proto.RegisterSignalExchangeServer(grpcServer, srv) grpcRootHandler := grpcHandlerFunc(grpcServer, metricsServer.Meter) diff --git a/signal/server/signal.go b/signal/server/signal.go index 7edbb4d34..f991b5d81 100644 --- a/signal/server/signal.go +++ b/signal/server/signal.go @@ -17,6 +17,8 @@ import ( "github.com/netbirdio/signal-dispatcher/dispatcher" + "github.com/netbirdio/netbird/shared/lifecycle" + "github.com/netbirdio/netbird/shared/profiling" "github.com/netbirdio/netbird/shared/signal/proto" "github.com/netbirdio/netbird/signal/metrics" "github.com/netbirdio/netbird/signal/peer" @@ -43,6 +45,8 @@ const ( labelRegistrationNotFound = "not_found" sendTimeout = 10 * time.Second + + applicationName = "signal" ) var ( @@ -51,6 +55,7 @@ var ( // Server an instance of a Signal server type Server struct { + lifecycle.StopHandlers registry *peer.Registry proto.UnimplementedSignalExchangeServer dispatcher *dispatcher.Dispatcher @@ -88,9 +93,17 @@ func NewServer(ctx context.Context, meter metric.Meter, metricsPrefix ...string) sendTimeout: sTimeout, } + stopProfiling := profiling.Start(applicationName) + s.OnStop(stopProfiling) + return s, nil } +// Stop runs the handlers registered with OnStop. +func (s *Server) Stop() { + s.RunStopHandlers() +} + // Send forwards a message to the signal peer func (s *Server) Send(ctx context.Context, msg *proto.EncryptedMessage) (*proto.EncryptedMessage, error) { log.Tracef("received a new message to send from peer [%s] to peer [%s]", msg.Key, msg.RemoteKey) diff --git a/upload-server/Dockerfile b/upload-server/Dockerfile index 3713d6f2a..8098a0186 100644 --- a/upload-server/Dockerfile +++ b/upload-server/Dockerfile @@ -1,4 +1,51 @@ -FROM gcr.io/distroless/base:debug -ENTRYPOINT [ "/go/bin/netbird-upload" ] -ARG TARGETPLATFORM -COPY ${TARGETPLATFORM}/netbird-upload /go/bin/netbird-upload +# syntax=docker/dockerfile:1 + +# Builds the upload server from source. Run it from the repository root: +# +# docker build -f upload-server/Dockerfile . +# +# Releases package the goreleaser-built binary with Dockerfile.release instead, +# which keeps the published image as it was (distroless base, running as root). +# +# The image runs as the base image's nonroot user (uid 65532), which owns the +# default STORE_DIR, /var/lib/netbird. A volume mounted there must be writable +# by that uid: a named Docker volume takes the directory's ownership on first +# use, a bind mount needs chown, and Kubernetes needs fsGroup: 65532. +# +# Build args: +# VARIANT=release|debug debug swaps the base for Chainguard busybox (a shell) +# VERSION stamped into the binary the same way goreleaser does +# +# Chainguard publishes only :latest for free, so the bases are pinned by digest +# and moved by Dependabot. + +ARG VARIANT=release + +# Pure Go: cross-compile from the build host instead of emulating the target. +FROM --platform=$BUILDPLATFORM golang:1.26.7-bookworm@sha256:e8c859f5632dcfde7b32d2012b4351728f6437930887c2f6a91ea242459e5514 AS builder +WORKDIR /app + +COPY go.mod go.sum ./ +RUN --mount=type=cache,target=/go/pkg/mod go mod download + +COPY . . +ARG TARGETOS +ARG TARGETARCH +ARG VERSION=development +RUN --mount=type=cache,target=/go/pkg/mod \ + --mount=type=cache,target=/root/.cache/go-build \ + CGO_ENABLED=0 GOOS=${TARGETOS} GOARCH=${TARGETARCH} go build -trimpath \ + -ldflags "-s -w -X github.com/netbirdio/netbird/version.version=${VERSION}" \ + -o /out/netbird-upload ./upload-server \ + && mkdir -p /out/var/lib/netbird + +FROM cgr.dev/chainguard/static:latest@sha256:41e17ed83c594a64a9396b6ab96dd26d5ddc290dacf4c177464712ff21ad534f AS base-release +FROM cgr.dev/chainguard/busybox:latest@sha256:b2953ab1cae4a6265e18cf675851bd99975211b150d7a014774911d76eb309ba AS base-debug + +# hadolint ignore=DL3006 +FROM base-${VARIANT} +COPY --from=builder --chown=65532:65532 /out/var/lib/netbird /var/lib/netbird +COPY --from=builder /out/netbird-upload /go/bin/netbird-upload +WORKDIR /var/lib/netbird +USER 65532:65532 +ENTRYPOINT ["/go/bin/netbird-upload"] diff --git a/upload-server/Dockerfile.release b/upload-server/Dockerfile.release new file mode 100644 index 000000000..3713d6f2a --- /dev/null +++ b/upload-server/Dockerfile.release @@ -0,0 +1,4 @@ +FROM gcr.io/distroless/base:debug +ENTRYPOINT [ "/go/bin/netbird-upload" ] +ARG TARGETPLATFORM +COPY ${TARGETPLATFORM}/netbird-upload /go/bin/netbird-upload diff --git a/upload-server/server/local.go b/upload-server/server/local.go index f7ca50011..859eeb9b2 100644 --- a/upload-server/server/local.go +++ b/upload-server/server/local.go @@ -8,9 +8,12 @@ import ( "os" "path/filepath" "strings" + "time" log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/shared/ratelimit" + "github.com/netbirdio/netbird/upload-server/types" ) @@ -20,11 +23,12 @@ const ( ) type local struct { - url string - dir string + url string + dir string + signer *signer } -func configureLocalHandlers(mux *http.ServeMux) error { +func configureLocalHandlers(mux *http.ServeMux, limiter *ratelimit.APIRateLimiter) error { envURL, ok := os.LookupEnv("SERVER_URL") if !ok { return fmt.Errorf("SERVER_URL environment variable is required") @@ -44,11 +48,17 @@ func configureLocalHandlers(mux *http.ServeMux) error { dir = envDir } - l := &local{ - url: envURL, - dir: dir, + uploadSigner, err := newSigner() + if err != nil { + return err } - mux.HandleFunc(types.GetURLPath, l.handlerGetUploadURL) + + l := &local{ + url: envURL, + dir: dir, + signer: uploadSigner, + } + mux.Handle(types.GetURLPath, limiter.Middleware(http.HandlerFunc(l.handlerGetUploadURL))) mux.HandleFunc(putURLPath+putHandler, l.handlePutRequest) return nil @@ -80,10 +90,11 @@ func (l *local) getUploadURL(objectKey string) (string, error) { return "", fmt.Errorf("failed to parse upload URL: %w", err) } newURL := parsedUploadURL.JoinPath(parsedUploadURL.Path, putURLPath, objectKey) + newURL.RawQuery = l.signer.sign(objectKey, time.Now()).Encode() return newURL.String(), nil } -const maxUploadSize = 150 << 20 +const maxUploadSize = 50 << 20 func (l *local) handlePutRequest(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPut { @@ -91,13 +102,6 @@ func (l *local) handlePutRequest(w http.ResponseWriter, r *http.Request) { return } - r.Body = http.MaxBytesReader(w, r.Body, maxUploadSize) - body, err := io.ReadAll(r.Body) - if err != nil { - http.Error(w, "request body too large or failed to read", http.StatusRequestEntityTooLarge) - return - } - uploadDir := r.PathValue("dir") if uploadDir == "" { http.Error(w, "missing dir path", http.StatusBadRequest) @@ -109,6 +113,19 @@ func (l *local) handlePutRequest(w http.ResponseWriter, r *http.Request) { return } + if err := l.signer.verify(uploadDir+"/"+uploadFile, r.URL.Query(), time.Now()); err != nil { + http.Error(w, "unauthorized", http.StatusUnauthorized) + log.Warnf("Rejected upload of %s/%s: %v", uploadDir, uploadFile, err) + return + } + + r.Body = http.MaxBytesReader(w, r.Body, maxUploadSize) + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, "request body too large or failed to read", http.StatusRequestEntityTooLarge) + return + } + cleanBase := filepath.Clean(l.dir) + string(filepath.Separator) dirPath := filepath.Clean(filepath.Join(l.dir, uploadDir)) @@ -125,14 +142,14 @@ func (l *local) handlePutRequest(w http.ResponseWriter, r *http.Request) { return } - if err = os.MkdirAll(dirPath, 0750); err != nil { + if err = os.MkdirAll(dirPath, 0o750); err != nil { http.Error(w, "failed to create upload dir", http.StatusInternalServerError) log.Errorf("Failed to create upload dir: %v", err) return } flags := os.O_WRONLY | os.O_CREATE | os.O_EXCL - f, err := os.OpenFile(filePath, flags, 0600) + f, err := os.OpenFile(filePath, flags, 0o600) if err != nil { if os.IsExist(err) { http.Error(w, "file already exists", http.StatusConflict) diff --git a/upload-server/server/local_test.go b/upload-server/server/local_test.go index 64b8fd228..3504087e9 100644 --- a/upload-server/server/local_test.go +++ b/upload-server/server/local_test.go @@ -8,19 +8,28 @@ import ( "os" "path/filepath" "testing" + "time" "github.com/stretchr/testify/require" "github.com/netbirdio/netbird/upload-server/types" ) +const testSigningKey = "test-signing-key-with-enough-length" + +func signedQuery(t *testing.T, objectKey string) string { + t.Helper() + s := &signer{key: []byte(testSigningKey)} + return s.sign(objectKey, time.Now()).Encode() +} + func Test_LocalHandlerGetUploadURL(t *testing.T) { mockURL := "http://localhost:8080" t.Setenv("SERVER_URL", mockURL) t.Setenv("STORE_DIR", t.TempDir()) mux := http.NewServeMux() - err := configureLocalHandlers(mux) + err := configureLocalHandlers(mux, newTestRateLimiter(t)) require.NoError(t, err) req := httptest.NewRequest(http.MethodGet, types.GetURLPath+"?id=test-file", nil) @@ -37,7 +46,6 @@ func Test_LocalHandlerGetUploadURL(t *testing.T) { require.Contains(t, response.URL, "test-file/") require.NotEmpty(t, response.Key) require.Contains(t, response.Key, "test-file/") - } func Test_LocalHandlePutRequest(t *testing.T) { @@ -45,13 +53,15 @@ func Test_LocalHandlePutRequest(t *testing.T) { mockURL := "http://localhost:8080" t.Setenv("SERVER_URL", mockURL) t.Setenv("STORE_DIR", mockDir) + t.Setenv(signingKeyVar, testSigningKey) mux := http.NewServeMux() - err := configureLocalHandlers(mux) + err := configureLocalHandlers(mux, newTestRateLimiter(t)) require.NoError(t, err) fileContent := []byte("test file content") - req := httptest.NewRequest(http.MethodPut, putURLPath+"/uploads/test.txt", bytes.NewReader(fileContent)) + req := httptest.NewRequest(http.MethodPut, + putURLPath+"/uploads/test.txt?"+signedQuery(t, "uploads/test.txt"), bytes.NewReader(fileContent)) rec := httptest.NewRecorder() mux.ServeHTTP(rec, req) @@ -69,13 +79,16 @@ func Test_LocalHandlePutRequest_PathTraversal(t *testing.T) { mockURL := "http://localhost:8080" t.Setenv("SERVER_URL", mockURL) t.Setenv("STORE_DIR", mockDir) + t.Setenv(signingKeyVar, testSigningKey) mux := http.NewServeMux() - err := configureLocalHandlers(mux) + err := configureLocalHandlers(mux, newTestRateLimiter(t)) require.NoError(t, err) fileContent := []byte("malicious content") - req := httptest.NewRequest(http.MethodPut, putURLPath+"/uploads/%2e%2e%2f%2e%2e%2fetc%2fpasswd", bytes.NewReader(fileContent)) + req := httptest.NewRequest(http.MethodPut, + putURLPath+"/uploads/%2e%2e%2f%2e%2e%2fetc%2fpasswd?"+signedQuery(t, "uploads/../../etc/passwd"), + bytes.NewReader(fileContent)) rec := httptest.NewRecorder() mux.ServeHTTP(rec, req) @@ -90,11 +103,13 @@ func Test_LocalHandlePutRequest_DirTraversal(t *testing.T) { mockDir := t.TempDir() t.Setenv("SERVER_URL", "http://localhost:8080") t.Setenv("STORE_DIR", mockDir) + t.Setenv(signingKeyVar, testSigningKey) - l := &local{url: "http://localhost:8080", dir: mockDir} + l := &local{url: "http://localhost:8080", dir: mockDir, signer: &signer{key: []byte(testSigningKey)}} body := bytes.NewReader([]byte("bad")) - req := httptest.NewRequest(http.MethodPut, putURLPath+"/x/evil.txt", body) + req := httptest.NewRequest(http.MethodPut, + putURLPath+"/x/evil.txt?"+signedQuery(t, "../../../tmp/evil.txt"), body) req.SetPathValue("dir", "../../../tmp") req.SetPathValue("file", "evil.txt") @@ -111,17 +126,20 @@ func Test_LocalHandlePutRequest_DuplicateFile(t *testing.T) { mockDir := t.TempDir() t.Setenv("SERVER_URL", "http://localhost:8080") t.Setenv("STORE_DIR", mockDir) + t.Setenv(signingKeyVar, testSigningKey) mux := http.NewServeMux() - err := configureLocalHandlers(mux) + err := configureLocalHandlers(mux, newTestRateLimiter(t)) require.NoError(t, err) - req := httptest.NewRequest(http.MethodPut, putURLPath+"/dir/dup.txt", bytes.NewReader([]byte("first"))) + req := httptest.NewRequest(http.MethodPut, + putURLPath+"/dir/dup.txt?"+signedQuery(t, "dir/dup.txt"), bytes.NewReader([]byte("first"))) rec := httptest.NewRecorder() mux.ServeHTTP(rec, req) require.Equal(t, http.StatusOK, rec.Code) - req = httptest.NewRequest(http.MethodPut, putURLPath+"/dir/dup.txt", bytes.NewReader([]byte("second"))) + req = httptest.NewRequest(http.MethodPut, + putURLPath+"/dir/dup.txt?"+signedQuery(t, "dir/dup.txt"), bytes.NewReader([]byte("second"))) rec = httptest.NewRecorder() mux.ServeHTTP(rec, req) require.Equal(t, http.StatusConflict, rec.Code) @@ -135,13 +153,15 @@ func Test_LocalHandlePutRequest_BodyTooLarge(t *testing.T) { mockDir := t.TempDir() t.Setenv("SERVER_URL", "http://localhost:8080") t.Setenv("STORE_DIR", mockDir) + t.Setenv(signingKeyVar, testSigningKey) mux := http.NewServeMux() - err := configureLocalHandlers(mux) + err := configureLocalHandlers(mux, newTestRateLimiter(t)) require.NoError(t, err) largeBody := make([]byte, maxUploadSize+1) - req := httptest.NewRequest(http.MethodPut, putURLPath+"/dir/big.txt", bytes.NewReader(largeBody)) + req := httptest.NewRequest(http.MethodPut, + putURLPath+"/dir/big.txt?"+signedQuery(t, "dir/big.txt"), bytes.NewReader(largeBody)) rec := httptest.NewRecorder() mux.ServeHTTP(rec, req) diff --git a/upload-server/server/ratelimit.go b/upload-server/server/ratelimit.go new file mode 100644 index 000000000..0a38e5cc3 --- /dev/null +++ b/upload-server/server/ratelimit.go @@ -0,0 +1,31 @@ +package server + +import ( + "os" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/shared/ratelimit" +) + +const defaultUploadBurst = 100 + +func newRateLimiter() *ratelimit.APIRateLimiter { + cfg, enabled := ratelimit.RateLimiterConfigFromEnv() + if os.Getenv(ratelimit.RateLimitingBurstEnv) == "" { + cfg.Burst = defaultUploadBurst + } + + // Rate limiting is enabled by default unless explicitly disabled + if os.Getenv(ratelimit.RateLimitingEnabledEnv) == "" { + enabled = true + } + + limiter := ratelimit.NewAPIRateLimiter(cfg) + limiter.SetEnabled(enabled) + + log.Infof("Upload URL rate limiting: enabled=%t rate=%.0f/min burst=%d trusted_proxies=%q", + limiter.Enabled(), cfg.RequestsPerMinute, cfg.Burst, os.Getenv(ratelimit.RateLimitingTrustedProxiesEnv)) + + return limiter +} diff --git a/upload-server/server/ratelimit_test.go b/upload-server/server/ratelimit_test.go new file mode 100644 index 000000000..18b5a34fd --- /dev/null +++ b/upload-server/server/ratelimit_test.go @@ -0,0 +1,60 @@ +package server + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/ratelimit" + "github.com/netbirdio/netbird/upload-server/types" +) + +func newTestRateLimiter(t *testing.T) *ratelimit.APIRateLimiter { + t.Helper() + + limiter := newRateLimiter() + t.Cleanup(limiter.Stop) + + return limiter +} + +func getUploadURL(t *testing.T, mux *http.ServeMux) int { + t.Helper() + + req := httptest.NewRequest(http.MethodGet, types.GetURLPath+"?id=test-file", nil) + req.Header.Set(types.ClientHeader, types.ClientHeaderValue) + rec := httptest.NewRecorder() + mux.ServeHTTP(rec, req) + + return rec.Code +} + +func Test_GetUploadURLIsRateLimited(t *testing.T) { + t.Setenv(ratelimit.RateLimitingBurstEnv, "2") + t.Setenv(ratelimit.RateLimitingRPMEnv, "1") + mux, _ := newLocalMux(t) + + require.Equal(t, http.StatusOK, getUploadURL(t, mux)) + require.Equal(t, http.StatusOK, getUploadURL(t, mux)) + require.Equal(t, http.StatusTooManyRequests, getUploadURL(t, mux)) +} + +func Test_RateLimitingIsOnByDefault(t *testing.T) { + t.Setenv(ratelimit.RateLimitingEnabledEnv, "") + t.Setenv(ratelimit.RateLimitingBurstEnv, "1") + mux, _ := newLocalMux(t) + + require.Equal(t, http.StatusOK, getUploadURL(t, mux)) + require.Equal(t, http.StatusTooManyRequests, getUploadURL(t, mux)) +} + +func Test_RateLimitingCanBeDisabled(t *testing.T) { + t.Setenv(ratelimit.RateLimitingEnabledEnv, "false") + t.Setenv(ratelimit.RateLimitingBurstEnv, "1") + mux, _ := newLocalMux(t) + + require.Equal(t, http.StatusOK, getUploadURL(t, mux)) + require.Equal(t, http.StatusOK, getUploadURL(t, mux)) +} diff --git a/upload-server/server/s3.go b/upload-server/server/s3.go index c0976acb5..55046fb9b 100644 --- a/upload-server/server/s3.go +++ b/upload-server/server/s3.go @@ -12,6 +12,8 @@ import ( "github.com/aws/aws-sdk-go-v2/service/s3" log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/shared/ratelimit" + "github.com/netbirdio/netbird/upload-server/types" ) @@ -21,7 +23,7 @@ type sThree struct { presignClient *s3.PresignClient } -func configureS3Handlers(mux *http.ServeMux) error { +func configureS3Handlers(mux *http.ServeMux, limiter *ratelimit.APIRateLimiter) error { bucket := os.Getenv(bucketVar) region, ok := os.LookupEnv("AWS_REGION") if !ok { @@ -40,7 +42,7 @@ func configureS3Handlers(mux *http.ServeMux) error { bucket: bucket, presignClient: s3.NewPresignClient(client), } - mux.HandleFunc(types.GetURLPath, handler.handlerGetUploadURL) + mux.Handle(types.GetURLPath, limiter.Middleware(http.HandlerFunc(handler.handlerGetUploadURL))) return nil } diff --git a/upload-server/server/s3_test.go b/upload-server/server/s3_test.go index 110b1b780..6c946d8d0 100644 --- a/upload-server/server/s3_test.go +++ b/upload-server/server/s3_test.go @@ -29,7 +29,7 @@ func Test_S3HandlerGetUploadURL(t *testing.T) { ctx := context.Background() c, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ ContainerRequest: testcontainers.ContainerRequest{ - Image: "quay.io/minio/minio:RELEASE.2025-04-22T22-12-26Z", + Image: "pgsty/silo:RELEASE.2026-09-16T00-00-00Z", ExposedPorts: []string{"9000/tcp"}, Env: map[string]string{ "MINIO_ROOT_USER": "minioadmin", @@ -90,7 +90,7 @@ func Test_S3HandlerGetUploadURL(t *testing.T) { t.Setenv(bucketVar, bucketName) mux := http.NewServeMux() - err = configureS3Handlers(mux) + err = configureS3Handlers(mux, newTestRateLimiter(t)) require.NoError(t, err) req := httptest.NewRequest(http.MethodGet, types.GetURLPath+"?id=test-file", nil) diff --git a/upload-server/server/server.go b/upload-server/server/server.go index 29ef72732..cd9b25c2e 100644 --- a/upload-server/server/server.go +++ b/upload-server/server/server.go @@ -10,6 +10,7 @@ import ( "github.com/google/uuid" log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/shared/ratelimit" "github.com/netbirdio/netbird/upload-server/types" ) @@ -19,7 +20,8 @@ const ( ) type Server struct { - srv *http.Server + srv *http.Server + limiter *ratelimit.APIRateLimiter } func NewServer() *Server { @@ -29,7 +31,7 @@ func NewServer() *Server { address = "0.0.0.0:8080" } mux := http.NewServeMux() - err := configureMux(mux) + limiter, err := configureMux(mux) if err != nil { log.Fatalf("Failed to configure server: %v", err) } @@ -38,7 +40,8 @@ func NewServer() *Server { }) return &Server{ - srv: &http.Server{Addr: address, Handler: mux}, + srv: &http.Server{Addr: address, Handler: mux}, + limiter: limiter, } } @@ -48,6 +51,9 @@ func (s *Server) Start() error { } func (s *Server) Stop() error { + if s.limiter != nil { + s.limiter.Stop() + } if s.srv != nil { log.Infof("Stopping upload server on %s", s.srv.Addr) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) @@ -57,13 +63,14 @@ func (s *Server) Stop() error { return nil } -func configureMux(mux *http.ServeMux) error { +func configureMux(mux *http.ServeMux) (*ratelimit.APIRateLimiter, error) { + limiter := newRateLimiter() + _, ok := os.LookupEnv(bucketVar) if ok { - return configureS3Handlers(mux) - } else { - return configureLocalHandlers(mux) + return limiter, configureS3Handlers(mux, limiter) } + return limiter, configureLocalHandlers(mux, limiter) } func getObjectKey(w http.ResponseWriter, r *http.Request) string { diff --git a/upload-server/server/signing.go b/upload-server/server/signing.go new file mode 100644 index 000000000..86ee785b2 --- /dev/null +++ b/upload-server/server/signing.go @@ -0,0 +1,90 @@ +package server + +import ( + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "fmt" + "net/url" + "os" + "strconv" + "time" + + log "github.com/sirupsen/logrus" +) + +const ( + signingKeyVar = "NB_UPLOAD_SIGNING_KEY" + + // signatureTTL matches the expiry the S3 backend puts on its presigned URLs. + signatureTTL = 15 * time.Minute + + expiryParam = "exp" + signatureParam = "sig" + + minSigningKeyLen = 32 +) + +type signer struct { + key []byte +} + +func newSigner() (*signer, error) { + if env, ok := os.LookupEnv(signingKeyVar); ok { + if env == "" { + return nil, fmt.Errorf("%s is set but empty", signingKeyVar) + } + if len(env) < minSigningKeyLen { + return nil, fmt.Errorf("%s must be at least %d bytes", signingKeyVar, minSigningKeyLen) + } + return &signer{key: []byte(env)}, nil + } + + key := make([]byte, 32) + if _, err := rand.Read(key); err != nil { + return nil, fmt.Errorf("generate signing key: %w", err) + } + log.Infof("%s not set, generated an ephemeral upload signing key", signingKeyVar) + + return &signer{key: key}, nil +} + +// sign returns the query parameters that authorize an upload of objectKey. +func (s *signer) sign(objectKey string, now time.Time) url.Values { + exp := now.Add(signatureTTL).Unix() + + v := url.Values{} + v.Set(expiryParam, strconv.FormatInt(exp, 10)) + v.Set(signatureParam, hex.EncodeToString(s.signature(objectKey, exp))) + + return v +} + +// verify reports whether query carries a still-valid signature over objectKey. +func (s *signer) verify(objectKey string, query url.Values, now time.Time) error { + exp, err := strconv.ParseInt(query.Get(expiryParam), 10, 64) + if err != nil { + return fmt.Errorf("malformed %s parameter", expiryParam) + } + + got, err := hex.DecodeString(query.Get(signatureParam)) + if err != nil { + return fmt.Errorf("malformed %s parameter", signatureParam) + } + + if !hmac.Equal(got, s.signature(objectKey, exp)) { + return fmt.Errorf("signature mismatch") + } + if now.Unix() >= exp { + return fmt.Errorf("upload URL expired") + } + + return nil +} + +func (s *signer) signature(objectKey string, exp int64) []byte { + mac := hmac.New(sha256.New, s.key) + fmt.Fprintf(mac, "%s\n%d", objectKey, exp) + return mac.Sum(nil) +} diff --git a/upload-server/server/signing_test.go b/upload-server/server/signing_test.go new file mode 100644 index 000000000..2fdd4aa3f --- /dev/null +++ b/upload-server/server/signing_test.go @@ -0,0 +1,148 @@ +package server + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "net/url" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/upload-server/types" +) + +func newLocalMux(t *testing.T) (*http.ServeMux, string) { + t.Helper() + + mockDir := t.TempDir() + t.Setenv("SERVER_URL", "http://localhost:8080") + t.Setenv("STORE_DIR", mockDir) + t.Setenv(signingKeyVar, testSigningKey) + + mux := http.NewServeMux() + require.NoError(t, configureLocalHandlers(mux, newTestRateLimiter(t))) + + return mux, mockDir +} + +func Test_LocalUploadURLRoundTrip(t *testing.T) { + mux, mockDir := newLocalMux(t) + + getReq := httptest.NewRequest(http.MethodGet, types.GetURLPath+"?id=test-file", nil) + getReq.Header.Set(types.ClientHeader, types.ClientHeaderValue) + getRec := httptest.NewRecorder() + mux.ServeHTTP(getRec, getReq) + require.Equal(t, http.StatusOK, getRec.Code) + + var response types.GetURLResponse + require.NoError(t, json.Unmarshal(getRec.Body.Bytes(), &response)) + + minted, err := url.Parse(response.URL) + require.NoError(t, err) + require.NotEmpty(t, minted.Query().Get(signatureParam)) + + content := []byte("bundle") + putRec := httptest.NewRecorder() + mux.ServeHTTP(putRec, httptest.NewRequest(http.MethodPut, minted.RequestURI(), bytes.NewReader(content))) + require.Equal(t, http.StatusOK, putRec.Code) + + written, err := os.ReadFile(filepath.Join(mockDir, response.Key)) + require.NoError(t, err) + require.Equal(t, content, written) +} + +func Test_LocalHandlePutRequest_RejectsUnauthorized(t *testing.T) { + expired := &signer{key: []byte(testSigningKey)} + + tests := []struct { + name string + query string + }{ + { + name: "no signature", + query: "", + }, + { + name: "tampered signature", + query: "exp=99999999999&sig=deadbeef", + }, + { + name: "malformed signature", + query: "exp=99999999999&sig=not-hex", + }, + { + // A signature is only good for the key it was minted for, so a URL + // handed out for one bundle cannot be replayed against another. + name: "signature for a different object", + query: signedQuery(t, "dir/other.txt"), + }, + { + name: "signature expiring this second", + query: expired.sign("dir/file.txt", time.Now().Add(-signatureTTL)).Encode(), + }, + { + name: "expired signature", + query: expired.sign("dir/file.txt", time.Now().Add(-signatureTTL-time.Minute)).Encode(), + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + mux, mockDir := newLocalMux(t) + + target := putURLPath + "/dir/file.txt" + if tc.query != "" { + target += "?" + tc.query + } + + rec := httptest.NewRecorder() + mux.ServeHTTP(rec, httptest.NewRequest(http.MethodPut, target, bytes.NewReader([]byte("payload")))) + + require.Equal(t, http.StatusUnauthorized, rec.Code) + + _, err := os.Stat(filepath.Join(mockDir, "dir", "file.txt")) + require.True(t, os.IsNotExist(err), "unauthorized upload should not be written") + }) + } +} + +func Test_SignerRejectsForeignKey(t *testing.T) { + minted := (&signer{key: []byte("one key")}).sign("dir/file.txt", time.Now()) + + err := (&signer{key: []byte("another key")}).verify("dir/file.txt", minted, time.Now()) + require.Error(t, err) +} + +func Test_NewSignerGeneratesEphemeralKey(t *testing.T) { + // Registers the restore hook, then clears the value for this test only. + t.Setenv(signingKeyVar, "") + os.Unsetenv(signingKeyVar) + + first, err := newSigner() + require.NoError(t, err) + second, err := newSigner() + require.NoError(t, err) + + require.NotEqual(t, first.key, second.key) + require.Len(t, first.key, 32) +} + +func Test_NewSignerRejectsEmptyKey(t *testing.T) { + t.Setenv(signingKeyVar, "") + + _, err := newSigner() + require.Error(t, err) +} + +func Test_NewSignerRejectsShortKey(t *testing.T) { + t.Setenv(signingKeyVar, strings.Repeat("a", minSigningKeyLen-1)) + + _, err := newSigner() + require.Error(t, err) +} diff --git a/util/file.go b/util/file.go index 52eb91c0f..4eb2f3ece 100644 --- a/util/file.go +++ b/util/file.go @@ -162,7 +162,7 @@ func writeBytes(ctx context.Context, file string, configDir string, configFileNa return fmt.Errorf("after temp file: %w", ctx.Err()) } - if err = os.Rename(tempFileName, file); err != nil { + if err = renameFile(tempFileName, file); err != nil { return fmt.Errorf("move %s to %s: %w", tempFileName, file, err) } @@ -195,7 +195,7 @@ func openOrCreateFile(file string) (*os.File, error) { // ReadJson reads JSON config file and maps to a provided interface func ReadJson(file string, res interface{}) (interface{}, error) { - f, err := os.Open(file) + f, err := openRead(file) if err != nil { return nil, err } @@ -248,7 +248,7 @@ func ListFiles(dir, pattern string) ([]string, error) { func ReadJsonWithEnvSub(file string, res interface{}) (interface{}, error) { envVars := getEnvMap() - f, err := os.Open(file) + f, err := openRead(file) if err != nil { return nil, err } diff --git a/util/file_nonwindows.go b/util/file_nonwindows.go new file mode 100644 index 000000000..c1db10244 --- /dev/null +++ b/util/file_nonwindows.go @@ -0,0 +1,16 @@ +//go:build !windows + +package util + +import "os" + +// openRead opens path for reading. Only Windows needs more than this: there a +// plain open holds the file against the rename that replaces it. +func openRead(path string) (*os.File, error) { + return os.Open(path) +} + +// renameFile replaces newpath with oldpath. +func renameFile(oldpath, newpath string) error { + return os.Rename(oldpath, newpath) +} diff --git a/util/file_read_test.go b/util/file_read_test.go new file mode 100644 index 000000000..d7276798f --- /dev/null +++ b/util/file_read_test.go @@ -0,0 +1,58 @@ +package util + +import ( + "errors" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestReadJson_ReadsTheFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + require.NoError(t, os.WriteFile(path, []byte(`{"SomeField": 7}`), 0o600)) + + var got TestConfig + _, err := ReadJson(path, &got) + + require.NoError(t, err) + assert.Equal(t, 7, got.SomeField, "the decoded value") +} + +// Callers tell a missing file from a broken one so they can seed a default in +// its place. The Windows path opens through a root and rebuilds the error, so +// the mapping has to survive that. +func TestReadJson_MissingFileIsErrNotExist(t *testing.T) { + dir := t.TempDir() + + for _, tc := range []struct { + name string + path string + }{ + {"missing file", filepath.Join(dir, "absent.json")}, + {"missing directory", filepath.Join(dir, "absent", "absent.json")}, + } { + t.Run(tc.name, func(t *testing.T) { + var got TestConfig + _, err := ReadJson(tc.path, &got) + + require.Error(t, err) + assert.ErrorIs(t, err, os.ErrNotExist) + assert.Contains(t, err.Error(), tc.path, "the error names the file the caller asked for") + }) + } +} + +func TestReadJson_MalformedFileIsNotErrNotExist(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + require.NoError(t, os.WriteFile(path, []byte("{not json"), 0o600)) + + var got TestConfig + _, err := ReadJson(path, &got) + + require.Error(t, err) + assert.False(t, errors.Is(err, os.ErrNotExist), + "a file that is there but unreadable must not be seeded over: %v", err) +} diff --git a/util/file_windows.go b/util/file_windows.go new file mode 100644 index 000000000..6abf6e308 --- /dev/null +++ b/util/file_windows.go @@ -0,0 +1,79 @@ +package util + +import ( + "errors" + "io/fs" + "os" + "path/filepath" +) + +// openRead opens path for reading without holding it against a rename. +// +// os.Open does not set FILE_SHARE_DELETE on Windows, so you cannot rename an +// open file like on UNIX. This caused concurrency issues with active state +// config file. +// +// os.Root opens through NtCreateFile with delete sharing, which is the +// behaviour Unix has. +// https://cs.opensource.google/go/go/+/refs/tags/go1.27.1:src/os/root_windows.go;drc=a4f5d9bbdbdf42da7e2d7e976ac85753c4db5d75;l=176 +func openRead(path string) (*os.File, error) { + root, err := os.OpenRoot(filepath.Dir(path)) + if err != nil { + // Names the file the caller asked for, not the directory the root + // failed on, so a missing directory reads like a missing file. + return nil, pathError("open", path, err) + } + defer func() { _ = root.Close() }() + + // The file outlives the root: closing a Root closes the directory handle it + // holds, not the files opened through it. + f, err := root.Open(filepath.Base(path)) + if err != nil { + return nil, pathError("open", path, err) + } + return f, nil +} + +// renameFile replaces newpath with oldpath, including while something holds +// newpath open for reading. +// +// os.Root.Rename asks for POSIX semantics, which unlink the destination +// immediately and leave open handles reading the version they opened. +// https://cs.opensource.google/go/go/+/master:src/internal/syscall/windows/at_windows.go;drc=a4f5d9bbdbdf42da7e2d7e976ac85753c4db5d75;l=384 +func renameFile(oldpath, newpath string) error { + dir := filepath.Dir(newpath) + if filepath.Dir(oldpath) != dir { + return os.Rename(oldpath, newpath) + } + + root, err := os.OpenRoot(dir) + if err != nil { + return os.Rename(oldpath, newpath) + } + defer func() { _ = root.Close() }() + + if err := root.Rename(filepath.Base(oldpath), filepath.Base(newpath)); err != nil { + return linkError("rename", oldpath, newpath, err) + } + return nil +} + +// pathError restores the full path on an error from a root, which names the +// file by the base name it was opened with. +func pathError(op, path string, err error) error { + var perr *fs.PathError + if errors.As(err, &perr) { + err = perr.Err + } + return &fs.PathError{Op: op, Path: path, Err: err} +} + +// linkError does the same as pathError for a rename, which reports both files +// by their base names. +func linkError(op, oldpath, newpath string, err error) error { + var lerr *os.LinkError + if errors.As(err, &lerr) { + err = lerr.Err + } + return &os.LinkError{Op: op, Old: oldpath, New: newpath, Err: err} +} diff --git a/util/file_windows_test.go b/util/file_windows_test.go new file mode 100644 index 000000000..eb7ba344f --- /dev/null +++ b/util/file_windows_test.go @@ -0,0 +1,116 @@ +package util + +import ( + "context" + "io" + "os" + "path/filepath" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// seedReplace lays out a write as writeBytes leaves it: the destination that +// exists and the temp file that is to take its place. +func seedReplace(t *testing.T) (src, dst string) { + t.Helper() + dir := t.TempDir() + src = filepath.Join(dir, ".tmpstate.json") + dst = filepath.Join(dir, "state.json") + require.NoError(t, os.WriteFile(src, []byte(`{"SomeField": 2}`), 0o600)) + require.NoError(t, os.WriteFile(dst, []byte(`{"SomeField": 1}`), 0o600)) + return src, dst +} + +// The reader has to share the file for delete, or the rename cannot take +// delete access on it. Regression test. +func TestRenameFile_ReplacesAFileBeingRead(t *testing.T) { + t.Run("a reader that shares delete", func(t *testing.T) { + src, dst := seedReplace(t) + + f, err := openRead(dst) + require.NoError(t, err) + defer f.Close() + + require.Error(t, os.Rename(src, dst), + "delete sharing alone has to be too little, or this test proves nothing") + require.NoError(t, renameFile(src, dst), "POSIX semantics have to get the replace through") + + // The handle stays on the file it opened, so a read in flight finishes + // on that version instead of seeing the replacement. + held, err := io.ReadAll(f) + require.NoError(t, err) + assert.JSONEq(t, `{"SomeField": 1}`, string(held), "the version the reader opened") + + landed, err := os.ReadFile(dst) + require.NoError(t, err) + assert.JSONEq(t, `{"SomeField": 2}`, string(landed), "the version the writer put there") + }) + + t.Run("a reader that does not", func(t *testing.T) { + src, dst := seedReplace(t) + + f, err := os.Open(dst) + require.NoError(t, err) + defer f.Close() + + require.Error(t, renameFile(src, dst), + "a plain read still holds the file, and the caller is owed that error") + }) + + t.Run("no readers at all", func(t *testing.T) { + src, dst := seedReplace(t) + + require.NoError(t, renameFile(src, dst)) + + landed, err := os.ReadFile(dst) + require.NoError(t, err) + assert.JSONEq(t, `{"SomeField": 2}`, string(landed), "the destination holds what replaced it") + }) +} + +// A config rewritten while it is being read, which is the daemon reading the +// active profile against a profile switch writing it. +func TestReadJsonWriteJson_Concurrently(t *testing.T) { + path := filepath.Join(t.TempDir(), "state.json") + require.NoError(t, WriteJson(context.Background(), path, &TestConfig{SomeField: 1})) + + var wg sync.WaitGroup + errs := make(chan error, 128) + + for i := 0; i < 8; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for r := 0; r < 50; r++ { + var got TestConfig + if _, err := ReadJson(path, &got); err != nil { + errs <- err + return + } + } + }() + } + + for i := 0; i < 2; i++ { + wg.Add(1) + go func(writer int) { + defer wg.Done() + for r := 0; r < 50; r++ { + if err := WriteJson(context.Background(), path, &TestConfig{SomeField: writer}); err != nil { + errs <- err + return + } + } + }(i) + } + + wg.Wait() + close(errs) + + for err := range errs { + assert.NoError(t, err, "a read and a write of the same config must not collide") + } +}