diff --git a/.git-branches.toml b/.git-branches.toml index d1818090f..4c34d7928 100644 --- a/.git-branches.toml +++ b/.git-branches.toml @@ -3,7 +3,7 @@ [branches] main = "main" perennials = [] -perennial-regex = "" +perennial-regex = "^release-" [create] new-branch-type = "feature" diff --git a/.github/dependabot.yml b/.github/dependabot.yml index b78b1417a..647e04936 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -3,8 +3,8 @@ updates: - package-ecosystem: "github-actions" directory: "/" schedule: - interval: "daily" - open-pull-requests-limit: 15 + interval: "weekly" + open-pull-requests-limit: 3 groups: actions: patterns: @@ -22,9 +22,12 @@ updates: directories: - "/" schedule: - interval: "daily" + interval: "weekly" open-pull-requests-limit: 15 groups: + golang-x-packages: + patterns: + - "golang.org/x/*" aws-sdk: patterns: - "github.com/aws/aws-sdk-go-v2/*" diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 8e68054bd..9b796f262 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -2,6 +2,12 @@ ## Issue ticket number and link + + ## Stack @@ -12,7 +18,9 @@ - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) -- [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). +- [ ] I ran and tested this change locally — I did not rely on CI to find out whether it works +- [ ] This PR has a single purpose (not a fix + refactor + feature in one) +- [ ] This change is a trivial fix, **OR** it links an issue the NetBird team agreed on beforehand. Changes to the public API, gRPC protocols, functionality behavior, CLI / service flags, or new features always need that agreement first. See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). diff --git a/.github/workflows/agent-network-e2e.yml b/.github/workflows/agent-network-e2e.yml index d78e3bbd3..9501c5fba 100644 --- a/.github/workflows/agent-network-e2e.yml +++ b/.github/workflows/agent-network-e2e.yml @@ -5,6 +5,20 @@ on: schedule: - cron: "0 3 * * *" workflow_dispatch: + inputs: + bedrock_model: + description: >- + Bedrock inference-profile id to drive the matrix with, exactly as + AWS issues it. Leave empty for the Sonnet 4.6 default. + required: false + default: "" + test_pattern: + description: >- + Package pattern to run. Defaults to the whole suite; narrow it to one + package (e.g. ./e2e/agentnetwork/...) when a run only needs that + package's answer and not the sixteen minutes the container suite costs. + required: false + default: "./e2e/..." concurrency: group: ${{ github.workflow }}-${{ github.ref }} @@ -51,6 +65,9 @@ jobs: # token (and URL, for gateways) is unset, so partial coverage is fine. OPENAI_TOKEN: ${{ secrets.E2E_OPENAI_TOKEN }} ANTHROPIC_TOKEN: ${{ secrets.E2E_ANTHROPIC_TOKEN }} + # Moonshot AI platform key (platform.kimi.ai); drives both Kimi wire + # shapes (OpenAI /v1 and Anthropic /anthropic) through kimi_api. + KIMI_TOKEN: ${{ secrets.E2E_KIMI_TOKEN }} VERCEL_URL: ${{ secrets.E2E_VERCEL_URL }} VERCEL_TOKEN: ${{ secrets.E2E_VERCEL_TOKEN }} OPENROUTER_URL: ${{ secrets.E2E_OPENROUTER_URL }} @@ -59,10 +76,16 @@ jobs: CLOUDFLARE_TOKEN: ${{ secrets.E2E_CLOUDFLARE_TOKEN }} AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.E2E_AWS_BEARER_TOKEN_BEDROCK }} AWS_REGION: ${{ secrets.E2E_AWS_REGION }} + # Bedrock model override: dispatch input wins, then the repo variable, else the test default. + AWS_BEDROCK_MODEL: ${{ inputs.bedrock_model || vars.E2E_AWS_BEDROCK_MODEL }} # Vertex (Anthropic-on-Vertex): SA + project required; region defaults # to "global", model to a pinned claude snapshot. GOOGLE_VERTEX_SA_BASE64: ${{ secrets.E2E_GOOGLE_VERTEX_SA_BASE64 }} GOOGLE_VERTEX_PROJECT: ${{ secrets.E2E_GOOGLE_VERTEX_PROJECT }} GOOGLE_VERTEX_REGION: ${{ secrets.E2E_GOOGLE_VERTEX_REGION }} GOOGLE_VERTEX_MODEL: ${{ secrets.E2E_GOOGLE_VERTEX_MODEL }} - run: go test -tags e2e -timeout 40m -v ./e2e/... + # Read through an env var rather than interpolated into the run + # script: a dispatch input reaching a shell command directly is a + # script-injection seam, however trusted the dispatcher. + TEST_PATTERN: ${{ inputs.test_pattern || './e2e/...' }} + run: go test -tags e2e -timeout 40m -v "$TEST_PATTERN" diff --git a/.github/workflows/check-license-dependencies.yml b/.github/workflows/check-license-dependencies.yml index 17c9fdc8d..81d293e4f 100644 --- a/.github/workflows/check-license-dependencies.yml +++ b/.github/workflows/check-license-dependencies.yml @@ -2,7 +2,7 @@ name: Check License Dependencies on: push: - branches: [main] + branches: [main, "release-*"] paths: - "go.mod" - "go.sum" diff --git a/.github/workflows/frontend-ui.yml b/.github/workflows/frontend-ui.yml index bb5bb4528..014c5c2ae 100644 --- a/.github/workflows/frontend-ui.yml +++ b/.github/workflows/frontend-ui.yml @@ -10,6 +10,7 @@ on: push: branches: - main + - "release-*" paths: - "client/ui/frontend/**" - "client/ui/i18n/**" @@ -86,7 +87,7 @@ jobs: ${{ runner.os }}-pnpm- - name: Install dependencies - run: pnpm install --frozen-lockfile + run: pnpm install --frozen-lockfile --ignore-scripts - name: Generate Wails bindings run: pnpm run bindings diff --git a/.github/workflows/golang-test-darwin.yml b/.github/workflows/golang-test-darwin.yml index 420749a0e..c17d8e775 100644 --- a/.github/workflows/golang-test-darwin.yml +++ b/.github/workflows/golang-test-darwin.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: concurrency: diff --git a/.github/workflows/golang-test-freebsd.yml b/.github/workflows/golang-test-freebsd.yml index 9c795e783..65c39147a 100644 --- a/.github/workflows/golang-test-freebsd.yml +++ b/.github/workflows/golang-test-freebsd.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: concurrency: diff --git a/.github/workflows/golang-test-linux.yml b/.github/workflows/golang-test-linux.yml index 0af506bba..004b78b3e 100644 --- a/.github/workflows/golang-test-linux.yml +++ b/.github/workflows/golang-test-linux.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: concurrency: diff --git a/.github/workflows/golang-test-windows.yml b/.github/workflows/golang-test-windows.yml index 50a5ba4d6..fb7b745d2 100644 --- a/.github/workflows/golang-test-windows.yml +++ b/.github/workflows/golang-test-windows.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: env: diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index b444a9900..586e1235b 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -45,7 +45,7 @@ jobs: display_name: Linux name: ${{ matrix.display_name }} runs-on: ${{ matrix.os }} - timeout-minutes: 15 + timeout-minutes: 25 steps: - name: Checkout code uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 @@ -79,4 +79,4 @@ jobs: skip-cache: true skip-save-cache: true cache-invalidation-interval: 0 - args: --timeout=12m + args: --timeout=20m diff --git a/.github/workflows/install-script-test.yml b/.github/workflows/install-script-test.yml index 1514caedc..61709501c 100644 --- a/.github/workflows/install-script-test.yml +++ b/.github/workflows/install-script-test.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: paths: - "release_files/install.sh" diff --git a/.github/workflows/mobile-build-validation.yml b/.github/workflows/mobile-build-validation.yml deleted file mode 100644 index 44e912c73..000000000 --- a/.github/workflows/mobile-build-validation.yml +++ /dev/null @@ -1,71 +0,0 @@ -name: Mobile - -on: - push: - branches: - - main - 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 - 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: Setup Android SDK - uses: android-actions/setup-android@40fd30fb8d7440372e1316f5d1809ec01dcd3699 # v4.0.1 - with: - cmdline-tools-version: 8512546 - - name: Setup Java - uses: actions/setup-java@1bcf9fb12cf4aa7d266a90ae39939e61372fe520 - with: - java-version: "11" - distribution: "adopt" - - name: NDK Cache - id: ndk-cache - uses: actions/cache@2c8a9bd7457de244a408f35966fab2fb45fda9c8 # v6.0.0 - with: - path: /usr/local/lib/android/sdk/ndk - key: ndk-cache-23.1.7779620 - - name: Setup NDK - run: /usr/local/lib/android/sdk/cmdline-tools/7.0/bin/sdkmanager --install "ndk;23.1.7779620" - - name: install gomobile - run: go install golang.org/x/mobile/cmd/gomobile@v0.0.0-20251113184115-a159579294ab - - name: gomobile init - run: gomobile init - - name: build android netbird lib - run: PATH=$PATH:$(go env GOPATH) gomobile bind -o $GITHUB_WORKSPACE/netbird.aar -javapkg=io.netbird.gomobile -ldflags="-checklinkname=0 -X golang.zx2c4.com/wireguard/ipc.socketDirectory=/data/data/io.netbird.client/cache/wireguard -X github.com/netbirdio/netbird/version.version=buildtest" $GITHUB_WORKSPACE/client/android - env: - CGO_ENABLED: 0 - ANDROID_NDK_HOME: /usr/local/lib/android/sdk/ndk/23.1.7779620 - ios_build: - name: "iOS / Build" - runs-on: macos-latest - 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: install gomobile - run: go install golang.org/x/mobile/cmd/gomobile@v0.0.0-20251113184115-a159579294ab - - name: gomobile init - run: gomobile init - - name: build iOS netbird lib - run: PATH=$PATH:$(go env GOPATH) gomobile bind -target=ios -bundleid=io.netbird.framework -ldflags="-X github.com/netbirdio/netbird/version.version=buildtest" -o ./NetBirdSDK.xcframework ./client/ios/NetBirdSDK - env: - CGO_ENABLED: 0 diff --git a/.github/workflows/no-new-replace.yml b/.github/workflows/no-new-replace.yml new file mode 100644 index 000000000..b906ce450 --- /dev/null +++ b/.github/workflows/no-new-replace.yml @@ -0,0 +1,78 @@ +name: No New Replace Directives + +on: + pull_request: + paths: + - "go.mod" + +permissions: + contents: read + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }} + cancel-in-progress: true + +jobs: + check-replace-directives: + name: check-replace-directives + runs-on: ubuntu-latest + steps: + - name: Checkout code + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + persist-credentials: false + fetch-depth: 0 + + - name: Install Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version-file: go.mod + + - name: Compare replace directives against the base branch + env: + BASE_SHA: ${{ github.event.pull_request.base.sha }} + run: | + set -euo pipefail + + # A replace directive only applies when this module is the main + # module. Anything importing netbird as a library, the embedded + # clients among them, resolves the replaced path upstream instead and + # fails to build against whatever the replacement provides. Requiring + # a fork under its own module path avoids that; a replace does not. + # + # go.mod is parsed rather than diffed so that reordering, comments and + # single-line versus block syntax do not register as changes. + # + # Versions are part of the key because a replace can be scoped to one + # version of a module. Keyed on paths alone, retargeting such a + # directive at a different version would read as unchanged. + list_replaces() { + go mod edit -json "$1" \ + | jq -r ' + def ref: .Path + (if (.Version // "") == "" then "" else " " + .Version end); + (.Replace // [])[] | "\(.Old | ref) => \(.New | ref)" + ' \ + | sort + } + + git show "${BASE_SHA}:go.mod" > /tmp/base-go.mod + list_replaces /tmp/base-go.mod > /tmp/base-replaces + list_replaces go.mod > /tmp/head-replaces + + added=$(comm -13 /tmp/base-replaces /tmp/head-replaces) + if [ -n "$added" ]; then + echo "::error::This PR adds a replace directive to go.mod:" + echo "$added" | sed 's/^/ /' + echo "" + echo "A replace directive applies only to the main module, so it does not" + echo "reach anything that imports netbird as a library. Require the module" + echo "under a path you control instead, as done for github.com/netbirdio/go-nat." + exit 1 + fi + + removed=$(comm -23 /tmp/base-replaces /tmp/head-replaces) + if [ -n "$removed" ]; then + echo "This PR removes replace directives:" + echo "$removed" | sed 's/^/ /' + fi + echo "No new replace directives." diff --git a/.github/workflows/pr-title-check.yml b/.github/workflows/pr-title-check.yml index 67d65356c..24d81b50f 100644 --- a/.github/workflows/pr-title-check.yml +++ b/.github/workflows/pr-title-check.yml @@ -16,6 +16,8 @@ jobs: const allowedTags = [ 'management', 'client', + 'android', + 'ios', 'signal', 'proxy', 'relay', diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 76a5b36ff..4d1945451 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -6,6 +6,7 @@ on: - "v*" branches: - main + - "release-*" pull_request: env: @@ -254,15 +255,23 @@ jobs: id: tag_and_push_images if: | (github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name == github.repository) || - (github.event_name == 'push' && github.ref == 'refs/heads/main') + (github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/heads/release-'))) run: | set -euo pipefail + # $GITHUB_REF / $GITHUB_EVENT_NAME are read from the runner + # environment rather than substituted into this script with the + # workflow expression syntax: branch names may legally contain + # $(…), and interpolating github.ref would execute it. resolve_tags() { - if [[ "${{ github.event_name }}" == "pull_request" ]]; then + if [[ "$GITHUB_EVENT_NAME" == "pull_request" ]]; then echo "pr-${{ github.event.pull_request.number }}" - else + elif [[ "$GITHUB_REF" == "refs/heads/main" ]]; then echo "main sha-$(git rev-parse --short HEAD)" + else + # Release branches get an immutable sha-* tag only — the floating + # "main" tag must never move from a release branch. + echo "sha-$(git rev-parse --short HEAD)" fi } @@ -475,6 +484,132 @@ jobs: path: dist/ retention-days: 3 + release_ui_gtk3: + # Legacy GTK3/WebKit2GTK 4.1 UI build for distros without WebKitGTK 6.0 + # (Ubuntu 22.04, Debian 12, RHEL 9, Fedora <=39). Runs on ubuntu-22.04 so + # the binary links against the oldest supported glibc. + runs-on: ubuntu-22.04 + outputs: + release_ui_gtk3_artifact_url: ${{ steps.upload_release_ui_gtk3.outputs.artifact-url }} + steps: + - name: Checkout + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + fetch-depth: 0 # It is required for GoReleaser to work properly + persist-credentials: false + + - name: Parse semver string + id: semver_parser + uses: netbirdio/shared-actions/actions/parse-semver@be5df6047383da2236e02243cceb857d8567c27e # v0.0.2 + + - name: Set snapshot flag + if: ${{ !startsWith(github.ref, 'refs/tags/v') }} + run: | + echo "flags=--snapshot" >> $GITHUB_ENV + + - name: Set build vars + if: ${{ startsWith(github.ref, 'refs/tags/v') }} + run: | + if [[ "x-${{ steps.semver_parser.outputs.prerelease }}" == "x-" && "x-${{ github.repository }}" == "x-netbirdio/netbird" ]]; then + echo "x-${{ github.repository }}" + echo "x-${{ steps.semver_parser.outputs.prerelease }}" + echo "SKIP_PUBLISH=false" >> $GITHUB_ENV + else + echo "x-${{ github.repository }}" + echo "x-${{ steps.semver_parser.outputs.prerelease }}" + fi + + - name: Set up Go + uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0 + with: + go-version-file: "go.mod" + cache: false + - name: Cache Go modules + # Restore-only from the release_ui cache written by trusted runs; the + # module cache is identical (same go.sum) and stale build-cache + # entries just miss. + uses: actions/cache/restore@2c8a9bd7457de244a408f35966fab2fb45fda9c8 # v6.0.0 + with: + path: | + ~/go/pkg/mod + ~/.cache/go-build + key: ${{ runner.os }}-ui-go-releaser-${{ hashFiles('**/go.sum') }} + restore-keys: | + ${{ runner.os }}-ui-go-releaser- + + - name: Install modules + run: go mod tidy + + - name: check git status + run: git --no-pager diff --exit-code + + - name: Set up Node.js + uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 + with: + node-version: '22' + + - name: Set up pnpm + uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0 + with: + version: 11 + + - name: Install dependencies + run: sudo apt update && sudo apt install -y -q libgtk-3-dev libwebkit2gtk-4.1-dev + + - name: Decode GPG signing key + if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository + env: + GPG_RPM_PRIVATE_KEY: ${{ secrets.GPG_RPM_PRIVATE_KEY }} + run: | + echo "$GPG_RPM_PRIVATE_KEY" | base64 -d > /tmp/gpg-rpm-signing-key.asc + echo "GPG_RPM_KEY_FILE=/tmp/gpg-rpm-signing-key.asc" >> $GITHUB_ENV + + - name: Install wails3 CLI + # Version derived from go.mod so the binding generator always matches + # the wails runtime the binary links against. + # -tags gtk3: the CLI links the wails runtime's cgo packages, and the + # default tags request gtk4/webkitgtk-6.0 pkg-config entries that do + # not exist on ubuntu-22.04. + run: | + WAILS_VERSION=$(go list -m -f '{{.Version}}' github.com/wailsapp/wails/v3) + go install -tags gtk3 github.com/wailsapp/wails/v3/cmd/wails3@$WAILS_VERSION + + - name: Run GoReleaser + uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2 + with: + version: ${{ env.GORELEASER_VER }} + args: release --config .goreleaser_ui_gtk3.yaml --clean ${{ env.flags }} + env: + GITHUB_TOKEN: ${{ secrets.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_UI_RPM_GTK3_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }} + SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }} + - name: Verify RPM signatures + run: | + docker run --rm -v $(pwd)/dist:/dist fedora:41 bash -c ' + dnf install -y -q rpm-sign curl >/dev/null 2>&1 + curl -sSL https://pkgs.netbird.io/yum/repodata/repomd.xml.key -o /tmp/rpm-pub.key + rpm --import /tmp/rpm-pub.key + echo "=== Verifying RPM signatures ===" + for rpm_file in /dist/*.rpm; do + [ -f "$rpm_file" ] || continue + echo "--- $(basename $rpm_file) ---" + rpm -K "$rpm_file" + done + ' + - name: Clean up GPG key + if: always() + run: rm -f /tmp/gpg-rpm-signing-key.asc + - name: upload non tags for debug purposes + id: upload_release_ui_gtk3 + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a #v7.0.1 + with: + name: release-ui-gtk3 + path: dist/ + retention-days: 3 + release_ui_darwin: runs-on: macos-latest outputs: @@ -688,7 +823,7 @@ jobs: comment_release_artifacts: name: Comment release artifacts runs-on: ubuntu-latest - needs: [release, release_ui, release_ui_darwin] + needs: [release, release_ui, release_ui_gtk3, release_ui_darwin] if: ${{ always() && github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name == github.repository }} permissions: contents: read @@ -700,12 +835,14 @@ jobs: env: RELEASE_RESULT: ${{ needs.release.result }} RELEASE_UI_RESULT: ${{ needs.release_ui.result }} + RELEASE_UI_GTK3_RESULT: ${{ needs.release_ui_gtk3.result }} RELEASE_UI_DARWIN_RESULT: ${{ needs.release_ui_darwin.result }} RELEASE_ARTIFACT_URL: ${{ needs.release.outputs.release_artifact_url }} LINUX_PACKAGES_ARTIFACT_URL: ${{ needs.release.outputs.linux_packages_artifact_url }} WINDOWS_PACKAGES_ARTIFACT_URL: ${{ needs.release.outputs.windows_packages_artifact_url }} MACOS_PACKAGES_ARTIFACT_URL: ${{ needs.release.outputs.macos_packages_artifact_url }} RELEASE_UI_ARTIFACT_URL: ${{ needs.release_ui.outputs.release_ui_artifact_url }} + RELEASE_UI_GTK3_ARTIFACT_URL: ${{ needs.release_ui_gtk3.outputs.release_ui_gtk3_artifact_url }} RELEASE_UI_DARWIN_ARTIFACT_URL: ${{ needs.release_ui_darwin.outputs.release_ui_darwin_artifact_url }} GHCR_IMAGES_MARKDOWN: ${{ needs.release.outputs.ghcr_images }} with: @@ -728,6 +865,7 @@ jobs: ['Windows packages', process.env.WINDOWS_PACKAGES_ARTIFACT_URL, process.env.RELEASE_RESULT], ['macOS packages', process.env.MACOS_PACKAGES_ARTIFACT_URL, process.env.RELEASE_RESULT], ['UI artifacts', process.env.RELEASE_UI_ARTIFACT_URL, process.env.RELEASE_UI_RESULT], + ['UI GTK3 artifacts', process.env.RELEASE_UI_GTK3_ARTIFACT_URL, process.env.RELEASE_UI_GTK3_RESULT], ['UI macOS artifacts', process.env.RELEASE_UI_DARWIN_ARTIFACT_URL, process.env.RELEASE_UI_DARWIN_RESULT], ]; @@ -784,7 +922,7 @@ jobs: trigger_signer: runs-on: ubuntu-latest - needs: [release, release_ui, release_ui_darwin, test_windows_installer] + needs: [release, release_ui, release_ui_gtk3, release_ui_darwin, test_windows_installer] if: startsWith(github.ref, 'refs/tags/') steps: - name: Trigger binaries sign pipelines diff --git a/.github/workflows/sync-tag.yml b/.github/workflows/sync-tag.yml index d99f88b54..088e538d5 100644 --- a/.github/workflows/sync-tag.yml +++ b/.github/workflows/sync-tag.yml @@ -9,21 +9,9 @@ concurrency: group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }} cancel-in-progress: true -# Receiving workflows (cloud sync-tag, mobile bump-netbird) expect the short -# tag form (e.g. v0.30.0), not refs/tags/v0.30.0 — github.ref_name, not github.ref. +# The receiving bump-netbird workflows expect the short tag form +# (e.g. v0.30.0), not refs/tags/v0.30.0 — github.ref_name, not github.ref. jobs: - trigger_sync_tag: - runs-on: ubuntu-latest - steps: - - name: Trigger release tag sync - uses: benc-uk/workflow-dispatch@31e2b3319479a63f0ab15bf800eff9e913504e26 # v1.3.2 - with: - workflow: sync-tag.yml - ref: main - repo: ${{ secrets.UPSTREAM_REPO }} - token: ${{ secrets.NC_GITHUB_TOKEN }} - inputs: '{ "tag": "${{ github.ref_name }}" }' - trigger_android_bump: runs-on: ubuntu-latest if: github.event.created && !github.event.deleted && startsWith(github.ref, 'refs/tags/v') && !contains(github.ref_name, '-') diff --git a/.github/workflows/test-infrastructure-files.yml b/.github/workflows/test-infrastructure-files.yml index 0a4f2e371..1313379ee 100644 --- a/.github/workflows/test-infrastructure-files.yml +++ b/.github/workflows/test-infrastructure-files.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: paths: - "infrastructure_files/**" @@ -249,78 +250,44 @@ jobs: docker compose exec management ls -l /var/lib/netbird/ | grep -i GeoLite2-City_[0-9]*.mmdb docker compose exec management ls -l /var/lib/netbird/ | grep -i geonames_[0-9]*.db - test-getting-started-script: + test-legacy-getting-started-scripts: runs-on: ubuntu-latest steps: - - name: Install jq - run: sudo apt-get install -y jq - - name: Checkout code uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: persist-credentials: false - - name: run script with Zitadel PostgreSQL - run: NETBIRD_DOMAIN=use-ip bash -x infrastructure_files/getting-started-with-zitadel.sh - - - name: test Caddy file gen postgres - run: test -f Caddyfile - - - name: test docker-compose file gen postgres - run: test -f docker-compose.yml - - - name: test management.json file gen postgres - run: test -f management.json - - - name: test turnserver.conf file gen postgres + - name: Verify fresh-install session cookie key hardening run: | - set -x - test -f turnserver.conf - grep external-ip turnserver.conf + grep -Fxq ' SESSION_COOKIE_ENCRYPTION_KEY=$(openssl rand -base64 32)' infrastructure_files/getting-started.sh + grep -Fxq ' sessionCookieEncryptionKey: "$SESSION_COOKIE_ENCRYPTION_KEY"' infrastructure_files/getting-started.sh + grep -Fxq ' install -m 600 /dev/null config.yaml' infrastructure_files/getting-started.sh + grep -Fxq ' openssl rand -base64 32' infrastructure_files/getting-started-enterprise.sh + grep -Fxq ' NETBIRD_SESSION_COOKIE_ENCRYPTION_KEY=$(rand_b64_key)' infrastructure_files/getting-started-enterprise.sh + grep -Fxq ' sessionCookieEncryptionKey: "${NETBIRD_SESSION_COOKIE_ENCRYPTION_KEY}"' infrastructure_files/getting-started-enterprise.sh - - name: test zitadel.env file gen postgres - run: test -f zitadel.env - - - name: test dashboard.env file gen postgres - run: test -f dashboard.env - - - name: test relay.env file gen postgres - run: test -f relay.env - - - name: test zdb.env file gen postgres - run: test -f zdb.env - - - name: Postgres run cleanup + - name: Verify Dex retirement notice run: | - docker compose down --volumes --rmi all - rm -rf docker-compose.yml Caddyfile zitadel.env dashboard.env machinekey/zitadel-admin-sa.token turnserver.conf management.json zdb.env + if infrastructure_files/getting-started-with-dex.sh >stdout.txt 2>stderr.txt; then + echo "Expected the retired Dex installer to fail" + exit 1 + fi + test ! -s stdout.txt + grep -Fq "Dex support is not deprecated." stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-quickstart" stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/identity-providers/local" stderr.txt + grep -Fq "removed in NetBird v0.80" stderr.txt - - name: run script with Zitadel CockroachDB - run: bash -x infrastructure_files/getting-started-with-zitadel.sh - env: - NETBIRD_DOMAIN: use-ip - ZITADEL_DATABASE: cockroach - - - name: test Caddy file gen CockroachDB - run: test -f Caddyfile - - - name: test docker-compose file gen CockroachDB - run: test -f docker-compose.yml - - - name: test management.json file gen CockroachDB - run: test -f management.json - - - name: test turnserver.conf file gen CockroachDB + - name: Verify Zitadel retirement notice run: | - set -x - test -f turnserver.conf - grep external-ip turnserver.conf - - - name: test zitadel.env file gen CockroachDB - run: test -f zitadel.env - - - name: test dashboard.env file gen CockroachDB - run: test -f dashboard.env - - - name: test relay.env file gen CockroachDB - run: test -f relay.env + if bash infrastructure_files/getting-started-with-zitadel.sh >stdout.txt 2>stderr.txt; then + echo "Expected the retired Zitadel installer to fail" + exit 1 + fi + test ! -s stdout.txt + grep -Fq "Zitadel support and existing Zitadel deployments are not deprecated." stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-quickstart" stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/identity-providers/zitadel" stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-guide" stderr.txt + grep -Fq "removed in NetBird v0.80" stderr.txt diff --git a/.github/workflows/ui-translations.yml b/.github/workflows/ui-translations.yml new file mode 100644 index 000000000..7d3b12f2d --- /dev/null +++ b/.github/workflows/ui-translations.yml @@ -0,0 +1,42 @@ +name: UI Translations + +on: + pull_request: + paths: + - "client/ui/i18n/locales/**" + - "client/ui/i18n/check-translations.mjs" + - ".github/workflows/ui-translations.yml" + push: + branches: + - main + paths: + - "client/ui/i18n/locales/**" + - "client/ui/i18n/check-translations.mjs" + +permissions: + contents: read + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }} + cancel-in-progress: true + +jobs: + check-translations: + name: Check translation key parity + runs-on: ubuntu-latest + timeout-minutes: 5 + steps: + - name: Checkout repository + uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + + - name: Set up Node.js + uses: actions/setup-node@v4 + with: + node-version: "22" + + # English (en) is the source of truth for translation keys; every other + # locale declared in _index.json must carry the exact same key set. + - name: Check translation key parity + run: node client/ui/i18n/check-translations.mjs diff --git a/.github/workflows/wasm-build-validation.yml b/.github/workflows/wasm-build-validation.yml index e8a12cdaf..5f21472e5 100644 --- a/.github/workflows/wasm-build-validation.yml +++ b/.github/workflows/wasm-build-validation.yml @@ -4,6 +4,7 @@ on: push: branches: - main + - "release-*" pull_request: concurrency: diff --git a/.goreleaser.yaml b/.goreleaser.yaml index 0753b1012..c5d260376 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -273,8 +273,8 @@ dockers_v2: - netbirdio/netbird - ghcr.io/netbirdio/netbird tags: - - "v{{ .Version }}-rootless" - - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}" + - "{{ .Version }}-rootless" + - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}rootless-latest{{ end }}" dockerfile: client/Dockerfile-rootless extra_files: - client/netbird-entrypoint.sh @@ -468,6 +468,13 @@ checksum: - glob: ./infrastructure_files/migrate-to-enterprise.sh release: + # The signing pipeline (netbirdio/sign-pipelines, dispatched by + # trigger_signer) marks the release latest once the Windows and macOS + # artifacts are signed. Without this override goreleaser marks it latest + # at publish time, while those artifacts are still unsigned. + make_latest: false + # Mark x.y.z-rc.* and other prerelease tags as prereleases on GitHub. + prerelease: auto extra_files: - glob: ./infrastructure_files/getting-started-with-zitadel.sh - glob: ./release_files/install.sh diff --git a/.goreleaser_ui.yaml b/.goreleaser_ui.yaml index 197fcd440..1c5bc41ac 100644 --- a/.goreleaser_ui.yaml +++ b/.goreleaser_ui.yaml @@ -24,6 +24,8 @@ builds: 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: + - production - id: netbird-ui-windows-amd64 dir: client/ui @@ -39,6 +41,8 @@ builds: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser - -H windowsgui mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production - id: netbird-ui-windows-arm64 dir: client/ui @@ -55,6 +59,8 @@ builds: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser - -H windowsgui mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production archives: - id: linux-arch @@ -87,7 +93,10 @@ nfpms: - src: client/ui/build/appicon.png dst: /usr/share/pixmaps/netbird.png dependencies: - - netbird + - netbird (>= 0.75.0) + - libgtk-4-1 (>= 4.14) + - libwebkitgtk-6.0-4 + - xdg-utils - maintainer: Netbird description: Netbird client UI. @@ -108,7 +117,10 @@ nfpms: - src: client/ui/build/appicon.png dst: /usr/share/pixmaps/netbird.png dependencies: - - netbird + - netbird >= 0.75.0 + - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) + - (webkitgtk6.0 or libwebkitgtk-6_0-4) + - xdg-utils rpm: signature: @@ -132,3 +144,11 @@ uploads: target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }} username: dev@wiretrustee.com method: PUT + +release: + # Uploads into the release created by the main .goreleaser.yaml run. + # make_latest stays false everywhere: the signing pipeline + # (netbirdio/sign-pipelines) marks the release latest after the Windows + # and macOS artifacts are signed. + make_latest: false + prerelease: auto diff --git a/.goreleaser_ui_darwin.yaml b/.goreleaser_ui_darwin.yaml index 96e15371a..8ca0e8da6 100644 --- a/.goreleaser_ui_darwin.yaml +++ b/.goreleaser_ui_darwin.yaml @@ -29,6 +29,8 @@ builds: 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: + - production universal_binaries: - id: netbird-ui-darwin @@ -41,3 +43,11 @@ checksum: name_template: "{{ .ProjectName }}_darwin_checksums.txt" changelog: disable: true + +release: + # Uploads into the release created by the main .goreleaser.yaml run. + # make_latest stays false everywhere: the signing pipeline + # (netbirdio/sign-pipelines) marks the release latest after the Windows + # and macOS artifacts are signed. + make_latest: false + prerelease: auto diff --git a/.goreleaser_ui_gtk3.yaml b/.goreleaser_ui_gtk3.yaml new file mode 100644 index 000000000..a9b2ca650 --- /dev/null +++ b/.goreleaser_ui_gtk3.yaml @@ -0,0 +1,144 @@ +version: 2 +env: + - SKIP_PUBLISH={{ if index .Env "SKIP_PUBLISH" }}{{ .Env.SKIP_PUBLISH }}{{ else }}true{{ end }} +project_name: netbird-ui + +before: + hooks: + # Bindings are gitignored; regenerate before the frontend build so + # the @wailsio/runtime Vite plugin can resolve them (vite refuses to + # build without them). + # -f '-tags gtk3': the generator type-checks client/ui, whose cgo imports + # would otherwise resolve gtk4/webkitgtk-6.0 pkg-config entries that do + # not exist on ubuntu-22.04. + - sh -c 'cd client/ui && wails3 generate bindings -clean=true -ts -f "-tags gtk3"' + - sh -c 'cd client/ui/frontend && pnpm install --frozen-lockfile && pnpm build' + +builds: + # Legacy GTK3 / WebKit2GTK 4.1 build for distros without WebKitGTK 6.0 + # (Ubuntu 22.04, Debian 12, RHEL 9, Fedora <=39). The gtk3 tag flips the + # Wails Linux backend to the GTK3 stack and swaps our GTK4-only XEmbed + # tray host for the pure-Go stub (client/ui/xembed_host_gtk3_linux.go). + # Must be built on the oldest supported glibc (ubuntu-22.04 runner). + - id: netbird-ui-gtk3 + dir: client/ui + binary: netbird-ui + env: + - CGO_ENABLED=1 + 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: + - production + - gtk3 + +archives: + - id: linux-gtk3-arch + name_template: "{{ .ProjectName }}-linux-gtk3_{{ .Version }}_{{ .Os }}_{{ .Arch }}" + builds: + - netbird-ui-gtk3 + +nfpms: + # Mutually-exclusive alternative to the GTK4 netbird-ui package -- both + # ship the same /usr/bin/netbird-ui from the shared stable/yum repos, so + # this one carries its own name and conflicts with the GTK4 package. + - maintainer: Netbird + description: Netbird client UI. + homepage: https://netbird.io/ + license: BSD-3-Clause + vendor: NetBird + id: netbird_ui_deb_gtk3 + package_name: netbird-ui-gtk3 + file_name_template: "{{ .PackageName }}_{{ .Version }}_{{ .Os }}_{{ .Arch }}" + builds: + - netbird-ui-gtk3 + formats: + - deb + scripts: + postinstall: "release_files/ui-post-install.sh" + contents: + - src: client/ui/build/linux/netbird.desktop + dst: /usr/share/applications/org.wails.netbird.desktop + - src: client/ui/build/appicon.png + dst: /usr/share/pixmaps/netbird.png + conflicts: + - netbird-ui + replaces: + - netbird-ui + dependencies: + - netbird (>= 0.75.0) + - libgtk-3-0 + - libwebkit2gtk-4.1-0 + - xdg-utils + + - maintainer: Netbird + description: Netbird client UI. + homepage: https://netbird.io/ + license: BSD-3-Clause + vendor: NetBird + id: netbird_ui_rpm_gtk3 + package_name: netbird-ui-gtk3 + file_name_template: "{{ .PackageName }}_{{ .Version }}_{{ .Os }}_{{ .Arch }}" + builds: + - netbird-ui-gtk3 + formats: + - rpm + scripts: + postinstall: "release_files/ui-post-install.sh" + contents: + - src: client/ui/build/linux/netbird.desktop + dst: /usr/share/applications/org.wails.netbird.desktop + - src: client/ui/build/appicon.png + dst: /usr/share/pixmaps/netbird.png + # No `replaces` here: nfpm maps it to rpm Obsoletes, which would make + # dnf swap installed GTK4 netbird-ui packages for this one on upgrade. + conflicts: + - netbird-ui + dependencies: + - netbird >= 0.75.0 + - (gtk3 or libgtk-3-0) + - (webkit2gtk4.1 or libwebkit2gtk-4_1-0) + - xdg-utils + + rpm: + signature: + key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}' + +# The GTK4 UI job shares project_name, so the default checksum file name would +# collide with it on the shared GitHub release. +checksum: + name_template: "{{ .ProjectName }}_gtk3_checksums.txt" + +changelog: + disable: true + +uploads: + - name: debian + skip: "{{ .Env.SKIP_PUBLISH }}" + ids: + - netbird_ui_deb_gtk3 + mode: archive + target: https://pkgs.wiretrustee.com/debian/pool/{{ .ArtifactName }};deb.distribution=stable;deb.component=main;deb.architecture={{ if .Arm }}armhf{{ else }}{{ .Arch }}{{ end }};deb.package= + username: dev@wiretrustee.com + method: PUT + + - name: yum + skip: "{{ .Env.SKIP_PUBLISH }}" + ids: + - netbird_ui_rpm_gtk3 + mode: archive + target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }} + username: dev@wiretrustee.com + method: PUT + +release: + # Uploads into the release created by the main .goreleaser.yaml run. + # make_latest stays false everywhere: the signing pipeline + # (netbirdio/sign-pipelines) marks the release latest after the Windows + # and macOS artifacts are signed. + make_latest: false + prerelease: auto diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..5497acb15 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,783 @@ +# NetBird Agent Guidelines + +**NetBird** is an open source connectivity platform: a WireGuard®-based overlay +network with a control plane. The **agent** (`client/`) runs on user machines as +a privileged daemon and manages the WireGuard interface, routing, firewall, and +DNS. **Management** (`management/`) is the control plane and REST/gRPC API, +**Signal** (`signal/`) brokers peer handshakes, **Relay** (`relay/`) carries +traffic when a direct tunnel is impossible, and **Proxy** (`proxy/`) is the +identity-aware proxy behind Agent Network. + +This file applies to the whole repository, and is the single source of truth for +agent guidance here. `CLAUDE.md` is a one-line pointer to it — keep the guidance +in this file, not duplicated there. + +## Contents + +- [STOP and ask the user before](#stop-and-ask-the-user-before) +- [Quick reference](#quick-reference) +- [Structure](#structure) +- [Where to look](#where-to-look) +- [Security](#security) +- [Agent conventions](#agent-conventions) +- [Repo-wide principles](#repo-wide-principles) +- [Type safety](#type-safety) +- [Concurrency and lifecycle](#concurrency-and-lifecycle) +- [Error handling](#error-handling) +- [Comments](#comments) +- [Testing](#testing) +- [Pitfalls](#pitfalls) +- [Commits, PRs, releases](#commits-prs-releases) +- [After you push: CI and review bots](#after-you-push-ci-and-review-bots) +- [Discussion and support](#discussion-and-support) + +## STOP and ask the user before + +- **Opening a pull request for anything beyond a trivial fix, without an agreed + ticket.** Ask the user directly: *"Is there a discussion or issue for this + change?"* NetBird is discussion-first — community reports start in + [Discussions](https://github.com/netbirdio/netbird/discussions), DevRel + validates them, and only validated discussions become issues. A PR that + changes behavior with no linked issue may be closed on arrival. If there is no + ticket, offer to draft the discussion post **instead of** the PR, and wait for + the user's call. Only typos, broken links, documentation corrections, and + one-line fixes that already have an issue can skip this. +- **Designing in any high-risk area** (see + [CONTRIBUTING.md](CONTRIBUTING.md#high-risk-areas)): public API and OpenAPI + schema, gRPC protos, behavior existing deployments would notice after an + upgrade, peer connectivity (ICE, NAT traversal, relay selection, WireGuard® or + Rosenpass key handling), client system integration (routing, firewall, DNS, + interface), authentication and authorization, CLI or service flags, config + file format, daemon IPC, store schema and migrations, or a new feature. The + design gets agreed in the ticket before code is written. +- **Writing a store migration or changing a persisted model.** Migrations are + one-way in the field and both the GORM and pgx paths may need the change. +- **Hand-editing generated code.** `*.pb.go`, `*.gen.go`, and mocks are outputs. + Edit the source (`.proto`, `openapi.yml`) and rerun the matching + `generate.sh`. +- **Adding, removing, or bumping a dependency**, and never vendor a fork. +- **Weakening a security control** — authentication, authorization, certificate + verification, privilege dropping, or peer identity checks — even when it is + the fastest way to make a test pass. +- **Force-pushing to `main`**, force-pushing any branch that is already under + review, amending pushed commits, or bypassing hooks with `--no-verify`. + +## Quick reference + +```bash +# Build +go build ./... +cd client && CGO_ENABLED=0 go build . # agent +cd management && go build . # management service +cd signal && go build . # signal service + +# Verify (run before every push) +go fmt ./... +make lint # golangci-lint on files changed vs origin/main (also the pre-push hook) +make lint-all # full-repository lint, matches CI +make test-unit # host-safe unit tests, -tags devcert, no sudo +make test-privileged # privileged-tagged suite in a Docker container with NET_ADMIN +make setup-hooks # wire make lint into .githooks/pre-push + +# Narrow runs +go test ./client/internal/dns/... +go test -race -run TestPeerConn ./client/internal/peer/... +PRIV_RUN=TestNftablesManager PRIV_PKGS=./client/firewall/nftables/... make test-privileged + +# Code generation (never hand-edit the output) +./shared/management/http/api/generate.sh # REST types from openapi.yml +./shared/management/proto/generate.sh +./shared/signal/proto/generate.sh +./client/proto/generate.sh +./flow/proto/generate.sh + +# Run locally (lab only, never on a machine you rely on) +sudo ./client/netbird up --log-level debug --log-file console +sudo ./client/netbird down # teardown: restores routing, firewall, DNS +./signal/signal run --log-level debug --log-file console +./management/management management --log-level debug --log-file console --config ./management.json +``` + +`netbird up` needs root and rewrites the host's routing table, firewall rules, +DNS configuration, and WireGuard® interface. Run it only in a disposable test +environment (a VM, container, or throwaway host) that you can rebuild, never on +a workstation or server whose connectivity matters. Run `sudo netbird down` +before you stop working, before rebuilding the binary, and on every failure +path, so the host's networking state is restored instead of left half-applied. +See [Pitfalls](#pitfalls) for why cleanup on every exit path matters. + +## Structure + +```text +netbird/ +├── client/ NetBird agent +│ ├── cmd/ agent CLI +│ ├── internal/ agent business logic (engine, peer, dns, routemanager, ...) +│ ├── server/ daemon for background execution +│ ├── proto/ daemon gRPC protos +│ ├── iface/ WireGuard® interface management +│ ├── firewall/ nftables, iptables, pf, WFP, userspace backends +│ ├── ssh/ built-in SSH server and client +│ ├── ui/ desktop UI (Wails v3 + React) +│ ├── android/, ios/ mobile bindings +│ ├── wasm/ WebAssembly build +│ └── mdm/, system/ MDM policy, host information +├── management/ control plane +│ └── server/ account, peer, groups, networks, posture, permissions, +│ settings, store, http (REST), idp, integrations, migration +├── signal/ handshake broker (peer/, server/) +├── relay/ relay service (protocol/, server/, healthcheck/) +├── proxy/ identity-aware proxy (llm/, acme/, accesslog/, middleware/, tcp/, udp/) +├── agent-network/ Agent Network overview +├── shared/ imported by both agent and services +│ ├── management/ proto/, client/, http/api (OpenAPI + generated types) +│ ├── signal/ proto/, client/ +│ └── relay/, auth/, sshauth/, metrics/ +├── e2e/ end-to-end suites and harness +├── encryption/, dns/, route/, stun/, sharedsock/, util/, flow/ +├── infrastructure_files/ docker compose and getting-started templates +└── release_files/ files packaged into releases +``` + +## Where to look + +| Task | Location | +| --------------------------- | ------------------------------------------------------------ | +| REST API / OpenAPI | `shared/management/http/api/` + `management/server/http/` | +| Management gRPC protocol | `shared/management/proto/` | +| Signal protocol | `shared/signal/proto/` | +| Daemon IPC protocol | `client/proto/` | +| Peer connection and NAT | `client/internal/peer/` | +| Network map handling | `client/internal/engine.go`, `shared/management/networkmap/` | +| Routing | `client/internal/routemanager/`, `route/` | +| Firewall backends | `client/firewall/` | +| DNS | `client/internal/dns/`, `dns/` | +| WireGuard® interface | `client/iface/` | +| Persistence and migrations | `management/server/store/`, `management/server/migration/` | +| IdP integrations | `management/server/idp/` | +| Permissions model | `management/server/permissions/` | +| LLM routing / Agent Network | `proxy/internal/llm/`, `agent-network/` | +| End-to-end tests | `e2e/` | + +## Security + +### Never fail open + +When a security check — access control, an IP restriction, an auth decision — +hits an error such as an unparseable value, an unavailable lookup, or a state it +does not recognize, it must **deny**. Never skip the check or allow the request +through because the check itself failed, and make the `default` and unknown cases +of a security-related `switch` deny rather than fall through. + +### Daemon RPC input is untrusted + +The agent runs as root (LocalSystem on Windows), so a daemon RPC crosses a +privilege boundary: treat every field as untrusted input rather than as something +the UI or CLI validated on the way in. + +When you add or change an RPC, ask what the handler does with caller input while +running as root. If the answer touches a filesystem path, a URL or host, or a +privileged state change, it needs a gate **in the handler** — a check in the client +that normally calls it is not a check at all. + +- **A caller-supplied path the daemon opens.** Never `os.Open` it as root. + Constrain it, then open it *as the caller* with `ipcauth.OpenOwnedFile`, which + opens `O_NOFOLLOW`, requires a regular file, and refuses a file the caller does + not own — so a symlink or hardlink aimed at a root-only file is rejected. +- **A caller-supplied URL or host the daemon fetches.** Restrict the scheme and + allow only known hosts for unprivileged callers. Prefer a lexical host + allowlist plus TLS verification over "resolve the host, then reject private + IPs": the resolve-then-trust pattern has a DNS-rebinding race (public IP at + check time, attacker IP at connect time), while a name allowlist has no IP + check to race. Never accept `http://` where `https://` is expected. +- **A privileged state change** (SSH root login, management URL, deregistration) + gates on the caller identity from `ipcauth.CallerIdentity(ctx)`. + +Caller identity comes from the kernel — `SO_PEERCRED`, `LOCAL_PEERCRED`, or the +named-pipe client token — and never from an RPC field. When +`ipcauth.CallerIdentity` reports that it could not determine an identity, **deny**; +do not fall back to treating the caller as the transport peer. + +## Agent conventions + +### Three networking modes + +Where packets actually flow depends on the mode the agent is running in. The +three are not interchangeable, so establish which one a change applies to — and +what it should do in the other two — before you write it. + +- **kernel mode** (Linux only): in-kernel WireGuard®. The kernel handles both + peer-to-peer and routed traffic, and ACLs are iptables or nftables rules. The + client programs kernel facilities but never sees the traffic itself. +- **userspace mode** (wireguard-go with a TUN): wireguard-go runs in-process. The + kernel handles peer-to-peer traffic once it leaves the TUN, while routed traffic + — exit nodes and network routes — goes through the userspace forwarder, which + terminates the connection and re-establishes it over OS sockets. Used on + platforms without kernel WireGuard® or when the user opts out. +- **netstack mode**: wireguard-go in-process with no TUN and no kernel + networking. The forwarder does all routing by stitching userspace sockets, and + listeners such as the embedded SSH and DNS servers bind on a gVisor netstack. + Used where the process cannot create a TUN device, such as the embedded client + (`client/embed/`) and the WASM build. + +### The overlay interface is not "WireGuard" + +Do not put "WireGuard" in identifiers or comments unless the code is genuinely +coupled to WireGuard® specifically — a wireguard-go call, a handshake field, a +kernel WireGuard® netlink attribute. For the interface, the host, peers, or +traffic in general, say "the NetBird interface", "the interface", or "the overlay". +Most firewall, routing, and DNS code is transport-agnostic, so a WireGuard® +reference there is simply inaccurate and rots as the transports change. + +### IPv6 is a soft feature + +The IPv6 overlay is opt-in dual-stack, and capability can change at runtime. Treat +it as soft rather than a requirement: + +- Gate local v6 paths on the interface accessor (`wgIface.Address().HasIPv6()`), + not on raw state fields, and skip the v6 path when the host has no v6 rather + than returning an error. +- Treat an empty or unparseable peer v6 address as "no v6 for that peer" and skip + it, keeping the v4 path working. +- Never let a missing v6 break v4. Fail-closed is for security checks; a + capability mismatch skips the v6 work and carries on. + +### Environment variables + +Name the variable in a constant and parse booleans with `strconv.ParseBool` rather +than comparing strings inline, so an unexpected value is logged instead of +silently meaning false: + +```go +const EnvDisableFeature = "NB_DISABLE_FEATURE" + +func isDisabledByEnv() bool { + val := os.Getenv(EnvDisableFeature) + if val == "" { + return false + } + disabled, err := strconv.ParseBool(val) + if err != nil { + log.Warnf("failed to parse %s: %v", EnvDisableFeature, err) + return false + } + return disabled +} +``` + +### Validating against protocol specs + +When a change depends on what a protocol actually mandates, read the specification +text from the [IETF datatracker](https://datatracker.ietf.org/) rather than a +summary, and check that you have the current RFC — the widely cited one for a +protocol is often superseded. Cite the section, not just the document, so a +reviewer can jump straight to the rule. + +## Repo-wide principles + +1. **Run `go fmt` on every modified Go file.** Formatting is not optional. +2. **Zero unaddressed linter warnings.** Fix what `golangci-lint` reports on code + you touch, and delete imports, helpers, and parameters your refactor orphaned. + Exception: unused parameters in shared code may be consumed by builds outside + this repository — do not remove them, ask instead. +3. **Function comments are mandatory for exported functions**, written as full + sentences with a period, starting with the identifier name. +4. **Prefer private functions and constants.** Export only what a caller outside + the package genuinely needs. +5. **Early returns and guard clauses.** Handle errors and edge cases first + instead of nesting `if`/`else` chains. +6. **Split complex functions.** If a function trips a complexity warning, break + it into named helpers rather than silencing the warning. +7. **Avoid LLM-slop tells:** em dashes, hedging narration, restating the diff in + prose, trailing summaries. Defaults, not absolute bans. Applies to code, + comments, commit messages, and PR descriptions alike. +8. **Concurrency: do a two-pass race analysis after every change** that touches + shared state, including reads of existing maps and slices. Guard them with a + mutex (or an atomic or channel where that fits better), keep critical + sections short, and run `go test -race` on the touched packages. See + [Concurrency and lifecycle](#concurrency-and-lifecycle) for the failure modes + to check for. +9. **Cross-platform builds must keep working.** The agent targets Linux, macOS, + Windows, FreeBSD, Android, and iOS. When you add a platform-specific file, + add the counterpart or a build-tagged fallback for the others. +10. **Never hand-edit generated files.** Change the source and regenerate. +11. **Never log secrets** — private keys, setup keys, tokens, PAT values — and + keep peer IPs and hostnames out of logs above debug level. + +## Type safety + +**No bare primitives for domain concepts.** A `string` parameter for an account +ID next to a `string` parameter for a peer ID is two bugs waiting to happen, +because the compiler cannot catch the swap. Declare the type once and use it +throughout, converting only at the boundaries where data enters or leaves — +protobuf, gRPC, HTTP, an external library. + +```go +type ServiceID string +type AccountID string + +// Internal: typed all the way through +func (r *Router) RemoveRoute(host SNIHost, svcID ServiceID) { ... } + +// Proto boundary: convert once, on the way in and on the way out +svcID := ServiceID(mapping.GetId()) +req.ServiceId = string(svcID) +``` + +- **IP addresses are `netip.Addr`**, not `string` and not `net.IP`. Parse at the + boundary and pass the typed value inward. +- **Always `Unmap()`** after parsing an address, after converting from `net.IP`, + and after extracting one from `RemoteAddr()`. This normalizes a v4-mapped v6 + address (`::ffff:10.1.2.3`) to plain v4 so IPv4 rules match it. A stored or + compared mapped address silently fails to match those rules. +- **Ports are `uint16`** internally; use `int` only where a library forces it and + convert immediately. +- **Enums are a typed string with constants**, so the valid set is discoverable + and a typo fails to compile. +- **Map keys follow the same rule**, and must be a real type (`type ServiceID + string`) rather than an alias (`type serviceID = string`) — an alias silently + accepts bare strings. + +## Concurrency and lifecycle + +Beyond the mutex hygiene in the principles above, check for these failure +modes. + +- **Never read a struct field inside a goroutine** when another goroutine may nil + or reassign it. Pass the value as a parameter, or capture it into a local before + launching. This matters most when `Stop()` nils a field without waiting for the + goroutine to finish. + + ```go + go func(ifaceName string) { // good: passed in, cannot be nilled underneath + m.Start(ctx, ifaceName) + }(iface.Name()) + ``` + +- **Never wait on a channel while holding a lock the sender needs.** Copy what you + need out from under the lock, release it, then wait. + + ```go + func (m *Manager) Stop() { + m.mu.Lock() + cancel, done := m.cancel, m.done + m.mu.Unlock() + if cancel != nil { + cancel() + <-done + } + } + ``` + +- **`Stop`/`Close` must be idempotent** — guard on an already-stopped flag or a + nil cancel — and must release the state they guarded. Clear maps and caches; + a cancelled goroutine holding a live map still pins that memory. Note that a + nil map only panics on writes; reads and iteration behave like an empty map, + so where post-close use must be rejected, check the stopped flag explicitly. +- **Publish coupled state only after every fallible step succeeds.** When several + fields form an invariant, build them into locals and assign them to the receiver + at the end. Assigning as you go leaves the object half-initialized when a later + step fails, so a readiness predicate reports ready while a coupled field is nil. + If an earlier step already had an external side effect — a created chain, an + opened handle, an inserted rule — roll it back before returning the error. +- **Clean up what you own on constructor error paths.** Once a constructor has + started something, every later error path must undo it: cancel a goroutine and + wait for it to exit, stop a ticker, close a watcher. The object is never + returned, so its `Close` will never run. +- **A failed `Start` must undo everything it started.** When a component brings up + several subsystems in sequence — connection manager, watchers, routing, DNS, + flow, persisted state — a failure partway through has to tear down the ones + already running, not just close the handle the error came from. Put the + already-started guard *before* that teardown path, so a rejected second `Start` + cannot dismantle the one that is running. + +## Error handling + +Use single-assignment form when the error is only needed inside the `if`: + +```go +// Good +if err := someCall(); err != nil { + return fmt.Errorf("context: %w", err) +} + +// Bad - unnecessary split +err := someCall() +if err != nil { + return fmt.Errorf("context: %w", err) +} +``` + +Use multiple assignment when the value is needed after the block: + +```go +result, err := someCall() +if err != nil { + return fmt.Errorf("context: %w", err) +} +``` + +Add short, meaningful context, and **do not** start `fmt.Errorf` messages with +obvious words like "failed to" or "error": + +```go +// Good +return fmt.Errorf("parse remote address: %w", err) +return fmt.Errorf("listen on %s: %w", addr, err) + +// Bad +return fmt.Errorf("failed to parse remote address: %w", err) +return fmt.Errorf("error listening on %s: %w", addr, err) + +// "failed" is fine in log messages +log.Debugf("failed to parse remote address: %v", err) +``` + +Skip the wrapping when a function only extracts or delegates and the wrap would +add nothing: + +```go +func parseAddr(addr string) (string, int, error) { + host, portStr, err := net.SplitHostPort(addr) + if err != nil { + return "", 0, err + } + // ... +} +``` + +Log the errors you choose not to act on: + +- `log.Debugf()` for errors that do not affect program flow but help debugging. +- `log.Tracef()` for very verbose errors that would otherwise spam logs. +- **Never ignore** errors from writes, network sends, or critical cleanup. +- Close errors may be ignored for read-only operations; log them at debug for + writes. + +**Do not log and return the same error.** It gets reported twice, from two places, +and the second reader cannot tell whether it happened once or twice. Return it and +let the caller decide. The exception is an API handler that has already written a +response. Internal helpers return errors rather than logging and swallowing them. + +**Never return a typed nil as an error.** A nil `*MyError` stored in an `error` +interface is not nil, so `err != nil` is true and callers take the failure path on +success. Return the error only where it is actually set: + +```go +if _, err := conn.Write(buf); err != nil { // good + return err +} +return nil +``` + +**Accumulate with `multierror` when an operation should continue past individual +failures** — teardown, cleanup, or setup where partial success is acceptable. +`client/errors.FormatErrorOrNil` returns nil for an empty accumulator, so callers +still see a plain nil on full success: + +```go +func (m *Manager) Cleanup() error { + var merr *multierror.Error + for _, r := range m.resources { + if err := r.Close(); err != nil { + merr = multierror.Append(merr, fmt.Errorf("close %s: %w", r.Name, err)) + } + } + return nberrors.FormatErrorOrNil(merr) +} +``` + +| Scenario | Approach | Why | +| --------------------- | --------------------- | ----------------------------------------- | +| Cleanup / teardown | Accumulate | Clean up as much as possible | +| Setup with rollback | Abort on first error | Partial state is invalid; undo what stuck | +| Setup with partial OK | Accumulate | Degraded operation is still useful | + +## Comments + +Comment the **why**, never the **what**. Default to no comment, and add one only +when a hidden constraint or workaround would surprise a future reader. Never +reference the current task, PR, or your own changes in a comment. + +```go +// Bad - trailing comments explaining the obvious +defer localConn.Close() // Close the connection +if err != nil { // Check if error occurred + +// Good +defer localConn.Close() + +// Good - explains a non-obvious constraint +// Use incremental checksum update per RFC 1624 for performance. +checksum = updateChecksum(checksum, oldPort, newPort) +``` + +### Length budget + +Neither of these is linter-enforced, so they are conventions the surrounding code +mostly follows rather than hard limits: + +- **Around 90 characters per line.** Wrap the comment rather than running well past + it. +- **Roughly 250 characters per comment**, about three wrapped lines. Doc comments + on exported identifiers may exceed it when the API genuinely needs the + explanation; inline comments inside a function body rarely should. + +The budget is a smell detector, not a rule to game. Do not compress a needed +explanation into cryptic shorthand to fit — if a block of code needs more than +250 characters of prose, the code is doing too much. Fix the code: + +- **Extract a named function.** A well-named function replaces the comment: the + name says *what*, the body shows *how*, and the comment you no longer write + was the *what* anyway. Clean Code calls this "explain yourself in code". +- **Extract a named constant or predicate.** `if isExpiredSetupKey(key)` needs + no comment; `if key.ExpiresAt.Before(now) && !key.Revoked && key.UsageLimit > 0` + does. +- **Keep the surviving comment for the why** — the RFC, the kernel quirk, the + ordering constraint. That part is usually one or two lines. + +### Long switch and if/else chains + +A `switch` whose cases carry multi-line explanations is the usual place this +budget is breached, and the comment is a symptom. In order of preference: + +1. **Extract each case body into a named function.** The case becomes one line, + the name carries the meaning, and the switch reads as a table of contents. +2. **Replace the switch with a lookup table** — `map[Kind]handlerFunc` — when the + branches are uniform. Adding a case stops meaning editing a growing function. +3. **Replace conditional with polymorphism** when branches vary by type and the + same switch shape starts appearing in more than one place. Clean Code's rule + of thumb: tolerate a switch statement if it appears **once**, is buried in a + factory that returns an interface, and no other switch dispatches on the same + type. A second switch over the same enum is the signal to introduce the + interface. + +Do not restructure a switch purely to satisfy the budget when the cases are one +line each and self-evident — a flat, boring `switch` over an enum is fine and +needs no comments at all. + +Explanatory comments in tests are welcome — they document the scenario being set +up, and the 250-character budget does not apply to them. + +## Testing + +- **Unit tests** live beside the code as `_test.go`. `make test-unit` runs the + host-safe set with `-tags devcert` and no sudo. +- **Privileged tests** carry the `privileged` build tag and mutate host + networking. They run through `make test-privileged`, inside a Docker container + with `NET_ADMIN`. Never bypass that harness by running them directly on the + host. +- **End-to-end suites** live in `e2e/` with a shared harness. +- **Test real behavior, not API existence.** Assert on the observable end state + a consumer would see — bytes that arrived, the packet after translation, the + row after the write — not merely that a method exists or returns an error. +- **Avoid mocks for code we own.** Exercise the real store, manager, or + controller and assert what the caller actually receives. +- **`require` for setup and preconditions, `assert` for the conditions under + test.** Use `require` whenever a later line would panic or be meaningless + otherwise. +- **Message guidance:** optional for `NoError`/`Error`; always give context for + comparison, boolean, and collection assertions. +- **Reproduce a bug before fixing it.** Write the test, watch it fail *for the + reason you expect* — a test that fails for an unrelated reason proves nothing — + then apply the fix and confirm it passes. Add the thin surrounding cases while + you are there. +- **Use `t.Setenv`** rather than `os.Setenv` so the previous value is restored on + cleanup. To test the unset case, call `t.Setenv` first to register the restore, + then `os.Unsetenv`. +- **Prefer `t.Cleanup` over `defer`** in any test with parallel subtests: the + parent function returns, running its `defer`s, while parallel subtests are + still suspended. Sequential subtests finish inside `t.Run`, so `defer` is safe + there, but `t.Cleanup` works in both cases. +- **Explanatory comments in tests are welcome.** Describe the scenario being set + up; the comment budget below does not apply to them. + +```go +server, err := StartTestServer() +require.NoError(t, err, "Test server setup must succeed") +defer server.Close() + +result, err := client.DoOperation() +assert.NoError(t, err) +assert.Equal(t, expectedResult, result, "Result should match expected") +``` + +## Pitfalls + +- **The agent runs as root.** Anything touching routing, firewall, DNS, or the + interface can take a user's machine off the network. Prefer a reversible + change and make sure cleanup runs on every exit path. +- **Management has two account loaders** (GORM and pgx). Adding a relation to an + account often means updating both, or it silently comes back empty in + production. +- **`go test ./...` without `-tags devcert` skips tests** that need the + development certificate. Use `make test-unit`. +- **`make lint` only checks the diff against `origin/main`.** CI runs + `make lint-all`; run it too before pushing a large change. +- **Protos are consumed by released clients.** An old agent must keep working + against a new Management, so fields are added, never renumbered or removed. +- **Windows requires the wintun driver**, and the daemon serves a named pipe + (`npipe://netbird`) rather than loopback TCP. Loopback TCP carries no caller + identity, so privileged operations are refused over it. + +## Commits, PRs, releases + +- **PR titles must start with a bracketed tag.** Before you propose a title, + **read [`.github/workflows/pr-title-check.yml`](.github/workflows/pr-title-check.yml) + and take the allowed tags from the `allowedTags` array in that file.** It is + the only source of truth, it changes as components are added, and the check + runs on every title edit — a tag that is not in that array is a red build. Do + not rely on a list memorized from anywhere else, including this file. + + ```text + [client] Authorize daemon IPC callers by their local identity + [management,client] Add MDM policy support + ``` + + Multiple tags are comma-separated inside one pair of brackets. Match the tag + to the component you actually changed, not to the one you read the most. + +- **Use the repository's PR template.** Fill in + [`.github/pull_request_template.md`](.github/pull_request_template.md) rather + than replacing it with your own summary: describe the change, link the issue, + tick the checklist honestly (including "ran locally" and "single purpose"), + and complete the documentation section. Do not tick a box you have not + verified, and do not delete rows that do not apply — the docs gate in CI reads + that section and fails when it is missing. + +- **Keep the PR description short.** Under 1000 words on top of the template's + own text, and usually far less — a few paragraphs. Reviewers read the diff; + the description exists to explain what the diff cannot say for itself. This is + well below what an agent will produce by default, so cut before you post. + +- **Body: why before what.** Lead with the problem and the reason for this + approach, then the shape of the change. No bullet list of files changed, no + per-function walkthrough, no restating the diff in prose, no trailing summary + section, no self-congratulatory closing line. + +- **No `Co-Authored-By` or tool-attribution trailers in the PR description**, + and none in commits either. Contributors own their contributions. Whatever + tooling produced the diff, the person opening the PR is its author: they have + read every line, they can explain why it works, they can answer review + questions without going back to a model, and they are accountable for the + consequences of merging it. Do not add a trailer, footer, or description line + that spreads that ownership onto a tool. + +- **Commit subjects follow the same `[scope] Subject` convention.** Keep the + subject short, and use the body for why before what. No bullet lists of files + changed. + +- **Push review fixes as separate commits.** The PR is squashed on merge, so + there is no reason to rewrite history mid-review; many small commits make the + re-review readable. + +- **Do not force-push a branch that is under review.** A force-push detaches + existing review comments from the lines they were written against, destroys + the "changes since your last review" diff a reviewer relies on, and discards + the CI history that showed which commit broke what. Add commits instead — + including for fixups and reverts. Force-push only when there is no + alternative: a rebase to clear a genuine conflict, or removing a secret or a + large binary that was committed by mistake. When you must, ask the user first, + then say so in a PR comment so reviewers know their anchors moved. Never + force-push `main`, and never force-push a branch you do not own. + +- **One PR, one purpose.** Split refactors out of fixes and fixes out of + features. + +- **Keep the PR small.** Size is the single strongest predictor of how long a PR + waits. Aim for **under ~400 changed lines across under ~20 files**; past + roughly **1000 lines or 50 files** a community PR is likely to be sent back to + be split, or left unreviewed until it is. Large PRs from outside the core team + may be blocked outright when the size was never agreed in the ticket — + reviewing a sprawling change against a privileged networking daemon is a + security risk in itself, not just a time cost. + + Judge the size by hand-written code: exclude generated output, `go.sum`, + vendored files, and test fixtures from the estimate, but do not use their + presence to argue a 3000-line PR is small. + + When a change genuinely cannot be small — a protocol migration, a + cross-component rename — agree the split in the ticket **before** writing + code, and land it as a sequence of PRs that each build, test, and make sense + on their own. Propose that split to the user rather than opening one large PR + and hoping. + + Prefer GitHub's stacked pull requests for such a sequence, rather than + hand-managing base branches: open each PR against the branch below it instead of + `main`, so every PR's diff shows only its own change. Merging a layer retargets + the PRs above it, and branch protections and required checks on the base branch + still apply to each one. + +- **User-facing changes need a docs PR** in + [netbirdio/docs](https://github.com/netbirdio/docs), linked from the PR + description. + +## After you push: CI and review bots + +Opening the PR is not the end of the task. Watch the run, read what the bots +say, and drive the PR to green before you report the work as done. + +```bash +gh pr checks --watch # all checks, live +gh run view --log-failed # only the failing steps +gh pr view --comments # bot and human review comments +``` + +**Never report a change as finished while checks are pending or red**, and never +describe a red PR as passing. If you ran out of turn before CI finished, say +which checks were still running. + +### The checks + +- **Go tests** — `golang-test-{linux,darwin,windows,freebsd}.yml`, sharded per + component. A failure in a component you did not touch is usually a real + interaction, not noise; read the log before assuming flake. +- **golangci-lint** — `golangci-lint.yml` runs the full repository, while + `make lint` only checks your diff. A clean local lint does not guarantee green + CI on a large change. +- **PR Title Check** — `pr-title-check.yml`, see above. +- **Codecov** — uploaded from the Linux test workflow with per-component flags + (`unit,client`, `unit,management`, `unit,relay`, `unit,proxy`, `unit,signal`, + `integration,management`). Coverage on new code should not go backwards. Add + tests for the paths you introduced; do not adjust thresholds or exclude files + to clear the report. +- **CodeRabbit** — configured in [`.coderabbit.yaml`](.coderabbit.yaml): `chill` + profile, auto-review on every non-draft PR, TypeScript/JavaScript/SVG paths + filtered out. Chat auto-reply is on, so `@coderabbitai` in a comment reaches + it. +- **SonarCloud** — project `netbirdio_netbird`, quality gate on new code (bugs, + vulnerabilities, code smells, duplication, coverage). +- **Snyk** — dependency and code scanning. + +Sonar and Snyk report as GitHub App checks rather than workflows in this +repository, so their detail lives on the PR check, not in the Actions logs. + +### Handling bot findings + +- **Read every comment and act on it.** Either fix it, or reply with the reason + it does not apply. Do not bulk-resolve threads to clear the count, and do not + silently ignore a finding because the check is advisory. +- **Bots are frequently wrong here.** NetBird has privileged, platform-specific, + and concurrency-heavy code that static analysis reads poorly. A confident + CodeRabbit or Sonar comment can still be nonsense. Verify the claim against + the code before you change anything — never edit correct code just to silence + a bot. +- **Security findings get the opposite default.** For a Snyk or Sonar + vulnerability, or a CodeRabbit comment about authentication, authorization, + certificate verification, or key handling, assume it is real until you have + disproved it. Surface it to the user rather than dismissing it yourself. +- **A new vulnerable dependency is a stop.** Bumping or replacing dependencies + needs the user's decision, as above. +- **Never change a workflow, threshold, lint exclusion, or bot config to make a + check pass.** If a check is genuinely wrong, say so and let the user decide. +- **Do not paper over flakes with blind re-runs.** Identify the failure first. If + it is a known flake, name it; if you cannot tell, report it as unresolved + rather than re-running until it goes green. + +## Discussion and support + +- Discussions: +- Slack: +- Docs: +- Security: — never in public +- Contribution process: [CONTRIBUTING.md](CONTRIBUTING.md) diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 000000000..764f406be --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1 @@ +See [AGENTS.md](AGENTS.md) for the agent guidelines in this repository. diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 261083783..9dea37ec8 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -1,6 +1,6 @@ # Contributing to NetBird -Thanks for your interest in contributing to NetBird. +Thanks for your interest in contributing to NetBird. There are many ways that you can contribute: - Reporting issues @@ -10,12 +10,99 @@ There are many ways that you can contribute: If you haven't already, join our slack workspace [here](https://docs.netbird.io/slack-url), we would love to discuss topics that need community contribution and enhancements to existing features. +## Ticket first, PR second + +**Open a ticket and wait for feedback before you open a pull request.** Every PR +that changes behavior must link to an issue the NetBird team has agreed on. A PR +that arrives without one may be closed and redirected to a discussion, no matter +how good the code is. + +Issues in this repository are maintainer-curated work items, so the flow starts +in [Discussions](https://github.com/netbirdio/netbird/discussions): + +1. **Open a discussion.** Use + [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) + for a bug, regression, or unexpected behavior, and + [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) + for a feature, enhancement, or integration idea. Setup and usage questions + belong in + [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support). + Never report a security vulnerability in public — follow the + [security policy](https://github.com/netbirdio/netbird/security/policy) + instead. +2. **Wait for feedback.** DevRel validates and reproduces the report, and a + maintainer confirms the direction. We may ask for more detail or propose a + different approach. Validated discussions become issues. +3. **Then write the code**, following the approach agreed in the issue, and open + the PR linking that issue. + +Trivial fixes — a typo, a broken link, a documentation correction, or a one-line +fix that already has an issue — can go straight to a PR. Everything else starts +with a ticket. When in doubt, ask in the discussion or on +[Slack](https://docs.netbird.io/slack-url); an hour of conversation up front +regularly saves a week of rework. + +### High-risk areas + +These always need the design discussed and agreed in the issue **before** you +write code: + +- **Public API** — REST / management API, OpenAPI schema, dashboard-facing contracts +- **gRPC protocols** — management, signal, relay, and client daemon protos +- **Functionality behavior** — anything existing deployments would experience differently after an upgrade +- **Peer connectivity** — ICE and NAT traversal, relay selection, WireGuard® and Rosenpass key handling +- **Client system integration** — routing, firewall, DNS, and interface management +- **Authentication and authorization** — IdP integration, tokens, permissions, cryptography +- **CLI / service flags**, configuration file format, and daemon IPC +- **Store and database schema** — models and migrations +- **New features** + +These surfaces are NetBird's contract with operators, self-hosters, and +downstream integrators, and changes to them have compatibility, security, and +release-planning implications. Agreeing on the direction early lets the PR +review focus on implementation rather than design. + +Typical bug fixes, internal refactors, documentation updates, and tests do not +need a design discussion, but should still be tied to an issue so the work is +visible and nobody duplicates it. + +### Using AI coding agents + +We have no policy for or against using an AI agent to write NetBird code. That +choice is yours, and we are not going to interrogate anyone about their tools. + +What we do have is a lot of incoming contributions that were plainly drafted with +one, and enough experience reviewing them to see the same avoidable problems +again and again: no ticket behind the change, a diff far too large to review, a +description longer than the code it describes, an approach that was never going +to be accepted, and an author who cannot answer questions about their own PR. +None of that is caused by the tooling — it is what happens when a tool is pointed +at a repository whose expectations it has never been told. + +So rather than a rule, there is a guide. [AGENTS.md](AGENTS.md) restates the +expectations from this document in the form agents read automatically +(`CLAUDE.md` points to it), so pointing your tool at the repository is usually +enough. Among other things it tells the agent to ask you for the +discussion or issue before drafting a PR, to keep the change small and +single-purpose, to run the tests locally, to use this repository's PR template +and title tags, and to write a description a reviewer can get through. + +The guardrails are the point, and they are the same ones we apply to everyone: an +agreed ticket, a change you have actually run, a diff small enough to review with +care, and an author who can explain it. Whatever wrote the diff, you are its +author — you own every line you submit and the consequences of opening a PR with it. + +We may assess whether a contribution is maintainable and whether its merged code +aligns with our security standards and design expectations. + ## Contents - [Contributing to NetBird](#contributing-to-netbird) + - [Ticket first, PR second](#ticket-first-pr-second) + - [High-risk areas](#high-risk-areas) + - [Using AI coding agents](#using-ai-coding-agents) - [Contents](#contents) - [Code of conduct](#code-of-conduct) - - [Discuss changes with the NetBird team first](#discuss-changes-with-the-netbird-team-first) - [Directory structure](#directory-structure) - [Development setup](#development-setup) - [Requirements](#requirements) @@ -24,6 +111,8 @@ If you haven't already, join our slack workspace [here](https://docs.netbird.io/ - [Build and start](#build-and-start) - [Test suite](#test-suite) - [Checklist before submitting a PR](#checklist-before-submitting-a-pr) + - [When we close a PR](#when-we-close-a-pr) + - [Translations](#translations) - [Other project repositories](#other-project-repositories) - [Contributor License Agreement](#contributor-license-agreement) @@ -34,42 +123,66 @@ Conduct which can be found in the file [CODE_OF_CONDUCT.md](CODE_OF_CONDUCT.md). By participating, you are expected to uphold this code. Please report unacceptable behavior to community@netbird.io. -## Discuss changes with the NetBird team first - -Changes to the **public API**, **gRPC protocols**, **functionality behavior**, **CLI / service flags**, or **new features** should be discussed with the NetBird team before you start the work. These surfaces are part of NetBird's contract with operators, self-hosters, and downstream integrators, and changes to them have compatibility, security, and release-planning implications that benefit from an early conversation. - -Open an issue or reach out on [Slack](https://docs.netbird.io/slack-url) to talk through what you have in mind. We'll help shape the change, flag any constraints we know about, and confirm the direction so the PR review can focus on implementation rather than design. - -Typical bug fixes, internal refactors, documentation updates, and tests do not need pre-discussion — open the PR directly. - ## Directory structure -The NetBird project monorepo is organized to maintain most of its individual dependencies code within their directories, except for a few auxiliary or shared packages. +The NetBird project monorepo keeps most of each component's code within its own +directory, except for a few auxiliary or shared packages. Protocol definitions +and the client-side service clients live under [/shared](/shared), because both +the agent and the services import them. -The most important directories are: +**Agent** -- [/.github](/.github) - Github actions workflow files and issue templates - [/client](/client) - NetBird agent code -- [/client/cmd](/client/cmd) - NetBird agent cli code +- [/client/cmd](/client/cmd) - NetBird agent CLI code - [/client/internal](/client/internal) - NetBird agent business logic code -- [/client/proto](/client/proto) - NetBird agent daemon GRPC proto files - [/client/server](/client/server) - NetBird agent daemon code for background execution -- [/client/ui](/client/ui) - NetBird agent UI code -- [/encryption](/encryption) - Contain main encryption code for agent communication -- [/iface](/iface) - Wireguard® interface code -- [/infrastructure_files](/infrastructure_files) - Getting started files containing docker and template scripts +- [/client/proto](/client/proto) - NetBird agent daemon gRPC proto files +- [/client/iface](/client/iface) - WireGuard® interface code +- [/client/firewall](/client/firewall) - Platform firewall backends (nftables, iptables, pf, WFP, userspace) +- [/client/ssh](/client/ssh) - Built-in SSH server and client +- [/client/ui](/client/ui) - NetBird agent UI code (Wails v3 + React) +- [/client/android](/client/android), [/client/ios](/client/ios) - Mobile platform bindings +- [/client/wasm](/client/wasm) - WebAssembly build of the agent +- [/client/mdm](/client/mdm) - MDM-delivered policy handling +- [/client/system](/client/system) - Host and system information collection + +**Control plane services** + - [/management](/management) - Management service code -- [/management/client](/management/client) - Management service client code which is imported by the agent code -- [/management/proto](/management/proto) - Management service GRPC proto files - [/management/server](/management/server) - Management service server code - [/management/server/http](/management/server/http) - Management service REST API code +- [/management/server/store](/management/server/store) - Persistence layer and migrations - [/management/server/idp](/management/server/idp) - Management service IDP management code -- [/release_files](/release_files) - Files that goes into release packages +- [/management/server/peer](/management/server/peer), [/management/server/groups](/management/server/groups), [/management/server/networks](/management/server/networks), [/management/server/posture](/management/server/posture), [/management/server/permissions](/management/server/permissions) - Core domain packages - [/signal](/signal) - Signal service code -- [/signal/client](/signal/client) - Signal service client code which is imported by the agent code - [/signal/peer](/signal/peer) - Signal service peer message logic -- [/signal/proto](/signal/proto) - Signal service GRPC proto files - [/signal/server](/signal/server) - Signal service server code +- [/relay](/relay) - Relay service code +- [/relay/protocol](/relay/protocol) - Relay wire protocol +- [/proxy](/proxy) - Identity-aware proxy used by Agent Network (LLM routing, ACME, access logs) +- [/agent-network](/agent-network) - Agent Network overview and documentation +- [/upload-server](/upload-server) - Debug bundle upload service + +**Shared code** + +- [/shared/management/proto](/shared/management/proto) - Management service gRPC proto files +- [/shared/management/client](/shared/management/client) - Management service client code which is imported by the agent code +- [/shared/management/http/api](/shared/management/http/api) - OpenAPI specification and generated REST API types +- [/shared/signal/proto](/shared/signal/proto) - Signal service gRPC proto files +- [/shared/signal/client](/shared/signal/client) - Signal service client code which is imported by the agent code +- [/shared/relay](/shared/relay) - Relay client and shared relay types +- [/shared/auth](/shared/auth), [/shared/sshauth](/shared/sshauth) - Shared authentication primitives +- [/encryption](/encryption) - Contain main encryption code for agent communication +- [/dns](/dns), [/route](/route), [/stun](/stun), [/sharedsock](/sharedsock), [/util](/util) - Shared networking and utility primitives +- [/flow](/flow) - Flow event protocol shared by the agent and Management + +**Build, test, and packaging** + +- [/.github](/.github) - Github actions workflow files, issue templates, and the pull request template +- [/e2e](/e2e) - End-to-end test suites and harness +- [/infrastructure_files](/infrastructure_files) - Getting started files containing docker and template scripts +- [/release_files](/release_files) - Files that goes into release packages +- [/tools](/tools) - Development and maintenance tooling ## Development setup @@ -234,12 +347,22 @@ cd client/ui task dev ``` -Pass daemon flags after `--`: +Pass daemon flags after `--`, pointing the UI at the socket the daemon serves: ``` -task dev -- --daemon-addr=tcp://127.0.0.1:41731 +task dev -- --daemon-addr=unix:///var/run/netbird.sock # Linux, macOS +task dev -- --daemon-addr=npipe://netbird # Windows ``` +On Windows the daemon serves a named pipe (`npipe://netbird`). Which path that +ends up being depends on what the daemon may create: as a service or elevated it +serves `\\.\pipe\ProtectedPrefix\Administrators\netbird`, which no unprivileged +process can take from it, and otherwise it falls back to `\\.\pipe\netbird`. +Clients try both and check who owns the pipe before using the plain one. Avoid +`tcp://127.0.0.1:41731`: loopback TCP carries no caller identity, so the daemon +refuses the operations that require an administrator and you will not exercise +those paths. + Production build (frontend assets embedded into the binary, output in `client/ui/bin/`): ``` @@ -324,23 +447,183 @@ The installer `netbird-installer.exe` will be created in root directory. ### Test suite -The tests can be started via: +The host-safe unit tests run as a normal user and leave host networking +untouched: ``` -cd netbird -go test -exec sudo ./... +make test-unit ``` + +Tests that need root and mutate host networking (firewall, routing, interface +management) carry the `privileged` build tag and run inside a +`--privileged --cap-add=NET_ADMIN` Docker container: + +``` +make test-privileged +``` + +Narrow a privileged run with environment variables: + +``` +PRIV_RUN=TestNftablesManager PRIV_PKGS=./client/firewall/nftables/... make test-privileged +``` + +Single packages can be run directly, adding `-race` when the change touches +shared state: + +``` +go test -race ./client/internal/dns/... +``` + > On Windows use a powershell with administrator privileges ## Checklist before submitting a PR -As a critical network service and open-source project, we must enforce a few things before submitting the pull-requests: + +As a critical network service and open source project, we must enforce a few +things before submitting a pull request. The +[pull request template](/.github/pull_request_template.md) mirrors this list — +fill it in rather than deleting it. + +### Link the issue + +The PR description must link the agreed issue (or the validated discussion it +came from). See [Ticket first, PR second](#ticket-first-pr-second). + +### Run it locally + +**If you can't run it, you can't submit it.** Build the affected components and +exercise the change on a real setup — see [Build and start](#build-and-start). +"CI will tell me" is not acceptable for a VPN agent that runs as root on other +people's machines. + +### Green CI, and answer the bots + +We do not start reviewing while CI is red. Get the pipeline green first — a +failing build, lint, or test means the PR is not ready for review. + +Alongside the test workflows, your PR is reviewed by CodeRabbit and scanned by +SonarCloud, Snyk, and Codecov. Read what they report and either fix it or reply +with why it does not apply; please do not resolve the threads without a +response. They are not always right — this codebase has privileged, +platform-specific, and concurrency-heavy paths that static analysis reads poorly +— so push back when a finding is wrong rather than changing correct code to +silence it. Security and dependency findings are the exception: treat those as +real until shown otherwise. Do not edit workflows, thresholds, or scanner +configuration to make a check pass. + +### One PR, one purpose + +Bug fix, refactor, feature: separate PRs. Mixed PRs are slow to review, hard to +revert, and may be closed with a request to split them. + +### Keep it small + +Size is the strongest predictor of how long a PR waits for review. Aim for under +roughly 400 changed lines across under 20 files. Past about 1000 lines or 50 +files, expect to be asked to split the change — and large PRs from outside the +core team may be blocked until the scope has been agreed in a ticket. This is +not only about reviewer time: NetBird's agent runs as root on other people's +machines, and a sprawling diff cannot be reviewed with the care that deserves. + +Measure by hand-written code, excluding generated output, `go.sum`, and +fixtures. If a change genuinely cannot be small — a protocol migration, a +cross-component rename — agree the split in the issue before you start, and land +it as a series of PRs that each build and make sense on their own. + +### Avoid force-pushing during review + +Once a PR is open, push new commits instead of rewriting history. A force-push +detaches existing review comments from their lines, throws away the +"changes since your last review" diff, and loses the CI history that showed +which commit broke what. Since we squash on merge, there is nothing to gain from +a tidy branch history. + +Force-pushing is sometimes unavoidable — rebasing to clear a real conflict, or +removing a secret or large binary committed by mistake. When that happens, leave +a comment on the PR so reviewers know their anchors moved. + +### Quality checks + +Run these from the repository root before pushing: + +```shell +go fmt ./... +make lint # golangci-lint on files changed against origin/main +make lint-all # full-repository lint, matches CI +make test-unit # host-safe unit tests +``` + +`make setup-hooks` wires `make lint` into a pre-push hook so the fast lint runs +automatically. If your change touches privileged paths (firewall, routing, +interface management), also run `make test-privileged`, which executes the +`privileged`-tagged suite inside a Docker container with `NET_ADMIN`. + +### Code standards + - Keep functions as simple as possible, with a single purpose - Use private functions and constants where possible - Comment on any new public functions - Add unit tests for any new public function +- Comment the **why**, not the **what** — explain non-obvious decisions, invariants, and constraints, not the line below +- Keep comments within 90 characters per line and roughly 250 characters per comment; when a block needs more explanation than that, extract a named function instead of writing a longer comment (see [AGENTS.md](AGENTS.md#length-budget)) + +### PR title and commits + +PR titles must start with a bracketed tag, enforced by +[pr-title-check.yml](/.github/workflows/pr-title-check.yml): + +```text +[client] Authorize daemon IPC callers by their local identity +[management,client] Add MDM policy support +``` + +Use a comma-separated list inside a single pair of brackets when a change spans +components. The `allowedTags` array in +[pr-title-check.yml](/.github/workflows/pr-title-check.yml) is the source of +truth — at the time of writing it accepts `management`, `client`, `signal`, +`proxy`, `relay`, `misc`, `infrastructure`, `self-hosted`, and `doc`. + +Commit subjects follow the same convention — keep them short and put the +reasoning in the body, why before what, with no bullet list of files changed. + +Keep the PR description itself under 1000 words on top of the template text. +Reviewers read the diff; the description explains what the diff cannot. > When pushing fixes to the PR comments, please push as separate commits; we will squash the PR before merging, so there is no need to squash it before pushing it, and we are more than okay with 10-100 commits in a single PR. This helps review the fixes to the requested changes. +### Documentation + +User-facing changes need a matching PR in +[netbirdio/docs](https://github.com/netbirdio/docs); link it in the PR +description, or state why documentation is not needed. + +## When we close a PR + +We would rather redirect early than let a PR sit. We may close one if: + +- It changes behavior with no linked issue, or the approach was never agreed with a maintainer +- The change was clearly never run or tested locally +- CI has been red without a response +- It mixes unrelated purposes, or the purpose is not clear +- It is far too large to review and the scope was never agreed in a ticket +- The author cannot answer questions about their own change — including PRs that read as unreviewed model output, where review turns into a relay between the maintainer and an LLM. Tooling is fine; unreviewed output is not, you are responsible for the code you sign your name to +- There has been no activity for 14 days after we requested changes + +A closed PR is not a rejected idea. Take it back to the +[discussion](https://github.com/netbirdio/netbird/discussions), settle the +approach, and reopen the work from there. + +## Translations + +Desktop UI translations are not contributed through pull requests. Translate on +[Crowdin](https://crowdin.com/project/netbird) instead: no ticket needed, just +join the project and pick your language. Crowdin syncs with this repository and +opens the service PRs itself, so hand-edited locale files would conflict with +the next sync. Style, terminology, and review guidance live in +[client/ui/i18n/TRANSLATING.md](client/ui/i18n/TRANSLATING.md). To request a +language the project does not offer yet, ask on the Crowdin project page or in +a [discussion](https://github.com/netbirdio/netbird/discussions). + ## Other project repositories NetBird project is composed of 3 main repositories: diff --git a/README.md b/README.md index 40c6b9ed5..3bcb4a035 100644 --- a/README.md +++ b/README.md @@ -130,7 +130,7 @@ In November 2022, NetBird joined the [StartUpSecure program](https://www.forschu ![CISPA_Logo_BLACK_EN_RZ_RGB (1)](https://user-images.githubusercontent.com/700848/203091324-c6d311a0-22b5-4b05-a288-91cbc6cdcc46.png) ### Acknowledgements -We build on open-source technologies like [WireGuard®](https://www.wireguard.com/), [Pion ICE](https://github.com/pion/ice), and [Rosenpass](https://rosenpass.eu). We greatly appreciate the work these projects are doing, and we'd love it if you could support them too (e.g., by starring or contributing). +We build on open source technologies like [WireGuard®](https://www.wireguard.com/), [Pion ICE](https://github.com/pion/ice), and [Rosenpass](https://rosenpass.eu). We greatly appreciate the work these projects are doing, and we'd love it if you could support them too (e.g., by starring or contributing). ### Legal This repository is licensed under the BSD-3-Clause license, which applies to all parts of the repository except for the directories management/, signal/ and relay/. diff --git a/SECURITY.md b/SECURITY.md index 745c66e61..cbcc975ba 100644 --- a/SECURITY.md +++ b/SECURITY.md @@ -1,12 +1,70 @@ # Security Policy -NetBird's goal is to provide a secure network. If you find a vulnerability or bug, please report it by opening an issue [here](https://github.com/netbirdio/netbird/issues/new?assignees=&labels=&template=bug-issue-report.md&title=) or by contacting us by email. - -There has yet to be an official bug bounty program for the NetBird project. - -## Supported Versions -- We currently support only the latest version +NetBird's goal is to provide a secure network. The client runs as a privileged service on every machine it is installed on, +so we take reports about it seriously and we publish what we fix. ## Reporting a Vulnerability -Please report security issues to `security@netbird.io` +**Please do not open a public issue for a security vulnerability.** Public issues are visible to everyone, including before +a fix is available. + +Report security issues one of these two ways: + +- **GitHub private vulnerability reporting** — [open a private report](https://github.com/netbirdio/netbird/security/advisories/new) + on this repository. This is the preferred route: it keeps the discussion, the draft advisory, and the credit in one place. +- **Email** — `security@netbird.io`. + +If the finding affects NetBird Cloud or our hosted infrastructure rather than the open source code, email us rather than +filing a repository report. + +### What to include + +A report is easier to act on when it contains: + +- The affected component (client, management, signal, relay, dashboard) and the version or commit you tested +- The platform and configuration, where relevant — operating system, self-hosted or NetBird Cloud, container or host install +- What an attacker needs before they can exploit it: network position, an account, local access, a specific privilege level +- Steps to reproduce, and a proof of concept if you have one +- The impact you believe it has + +Partial reports are still welcome. If you are unsure whether something is a security issue, send it to `security@netbird.io` +and let us make that call. + +## What to expect from us + +- **We acknowledge your report** and tell you whether we can reproduce it. +- **We work with you on severity and scope.** If we assess it differently than you do, we will explain why rather than + silently downgrade it. +- **We fix and release**, then publish a [GitHub Security Advisory](https://github.com/netbirdio/netbird/security/advisories) + naming the affected version range and the patched version. +- **We credit reporters who want to be credited.** Tell us the name or handle you would like used, or that you would rather + stay anonymous. +- **We keep you in the loop** until the advisory is published. + +We ask that you give us a reasonable opportunity to ship a fix before disclosing the issue publicly, and that you avoid +accessing, modifying, or exfiltrating data belonging to other people while testing. Testing against your own installation +or your own account is always fine. + +## Supported Versions + +We support the latest release. Security fixes ship in the next version rather than as backports to older releases, so +upgrading to the current release is how you get them. + +Release notifications are available by watching [releases](https://github.com/netbirdio/netbird/releases). + +## Published advisories + +Every vulnerability we fix is published as a GitHub Security Advisory on the +[advisories page](https://github.com/netbirdio/netbird/security/advisories), including the affected version range, the +patched version, and the reporter's credit. Advisories for the Go module are also distributed through the Go vulnerability +database, so `govulncheck` will report them against your dependencies. + +## Bug bounty + +There is no official bug bounty program for the NetBird project. We credit reporters in advisories, and we are grateful for +the work, but we cannot currently offer payment for reports. + +## Non-security bugs + +For bugs that are not security issues, please use the +[issue tracker](https://github.com/netbirdio/netbird/discussions/new/choose). diff --git a/agent-network/README.md b/agent-network/README.md index 1997ea299..5211fe8f9 100644 --- a/agent-network/README.md +++ b/agent-network/README.md @@ -40,6 +40,35 @@ You can then use this private endpoint to configure your AI agents, whether that Full step-by-step setup: **https://docs.netbird.io/agent-network/quickstart** +## Client settings that don't follow the endpoint + +Most of an agent's traffic follows the base URL you hand it, but a few +client-side checks call their vendor directly and never reach the proxy. On a +network that blocks direct egress they fail even though inference works, so +they are worth setting once when you roll the endpoint out. + +For Claude Code: + +- **Fast mode** checks availability against `api.anthropic.com` rather than the + configured base URL. Set `CLAUDE_CODE_SKIP_FAST_MODE_ORG_CHECK=1` when the + agent authenticates with `ANTHROPIC_AUTH_TOKEN` alone (the usual shape when + the proxy injects the real provider key) or when a TLS-inspecting proxy + answers the check itself. Set + `CLAUDE_CODE_SKIP_FAST_MODE_NETWORK_ERRORS=1` when the network refuses the + connection outright. Fast mode is an Anthropic-API feature, so it is + unavailable on a Bedrock- or Vertex-backed endpoint whatever you set. +- **Model discovery** is off by default. Set + `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` for the picker to list the + models your policies authorise; the proxy filters the response to that set. + The client gives discovery a three-second budget and treats any redirect as + a failure, so the endpoint must serve `/v1/models` directly. +- **The WebFetch domain safety check** also calls `api.anthropic.com` directly + and is unaffected by the variables above. + +Allowing direct egress to `api.anthropic.com` covers the network cases but not +the credential one, where the check reaches Anthropic and is rejected because +the agent presents a proxy-issued key. + ## Architecture Agent Network is built on two existing NetBird capabilities: diff --git a/client/android/client.go b/client/android/client.go index 2266ff53d..5bd0d1e10 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -7,6 +7,7 @@ import ( "fmt" "os" "slices" + "strings" "sync" "time" @@ -14,6 +15,7 @@ import ( log "github.com/sirupsen/logrus" + nbAnonymize "github.com/netbirdio/netbird/client/anonymize" "github.com/netbirdio/netbird/client/iface/device" "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/debug" @@ -24,6 +26,7 @@ import ( "github.com/netbirdio/netbird/client/internal/routemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/net" + "github.com/netbirdio/netbird/client/netevents" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/formatter" "github.com/netbirdio/netbird/route" @@ -31,10 +34,12 @@ import ( types "github.com/netbirdio/netbird/upload-server/types" ) -// ConnectionListener export internal Listener for mobile -type ConnectionListener interface { - peer.Listener -} +// AnonymizeLevelDefault and AnonymizeLevelStrict are the accepted +// anonymizeLevel values for DebugBundle. +const ( + AnonymizeLevelDefault = nbAnonymize.LevelDefaultString + AnonymizeLevelStrict = nbAnonymize.LevelStrictString +) // TunAdapter export internal TunAdapter for mobile type TunAdapter interface { @@ -56,6 +61,12 @@ type DnsReadyListener interface { dns.ReadyListener } +// TunSettings is a snapshot of the settings the TUN device is rebuilt with +type TunSettings struct { + Routes string + SearchDomains string +} + func init() { formatter.SetLogcatFormatter(log.StandardLogger()) } @@ -70,18 +81,43 @@ type Client struct { deviceName string uiVersion string networkChangeListener listener.NetworkChangeListener + // netMgr outlives engine restarts: it mirrors the OS connectivity, not + // the engine lifecycle. Run and RunWithoutLogin inject its state and + // sweeper into each new ConnectClient. + netMgr *netevents.Manager stateMu sync.RWMutex connectClient *internal.ConnectClient config *profilemanager.Config cacheDir string + // Identifies the running profile for the SSO login hint; see profile_state.go. + cfgPath string + + stateChangeMu sync.Mutex + stateChangeSubID string + eventSub *peer.EventSubscription + // Closed to stop the watch goroutines from delivering buffered items to a + // listener that has been removed or replaced. See stopStateChangeWatchLocked. + stateChangeDone chan struct{} + + // Latched "the server wants an interactive login": survives the engine + // restarts that replace the run loop's context state. See Client.Status. + // Guarded by loginRequiredMu together with loginCleared, which counts + // clears so a stale observation cannot re-latch over one. + loginRequiredMu sync.Mutex + loginRequired bool + loginCleared uint64 + + extendMu sync.Mutex + extendCancel context.CancelFunc } -func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cc *internal.ConnectClient) { +func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cfgPath string, cc *internal.ConnectClient) { c.stateMu.Lock() defer c.stateMu.Unlock() c.config = cfg c.cacheDir = cacheDir + c.cfgPath = cfgPath c.connectClient = cc } @@ -91,6 +127,16 @@ func (c *Client) stateSnapshot() (*profilemanager.Config, string, *internal.Conn return c.config, c.cacheDir, c.connectClient } +// authSnapshot returns the config together with the path it was loaded from, in +// one lock: the path identifies the profile whose account email backs the login +// hint, so reading it separately could pair one profile's config with another's +// hint when a profile switch lands in between. +func (c *Client) authSnapshot() (*profilemanager.Config, string, *internal.ConnectClient) { + c.stateMu.RLock() + defer c.stateMu.RUnlock() + return c.config, c.cfgPath, c.connectClient +} + func (c *Client) getConnectClient() *internal.ConnectClient { c.stateMu.RLock() defer c.stateMu.RUnlock() @@ -102,14 +148,17 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd execWorkaround(androidSDKVersion) net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket) + system.SetIFaceDiscover(iFaceDiscover) + recorder := peer.NewRecorder("") return &Client{ deviceName: deviceName, uiVersion: uiVersion, tunAdapter: tunAdapter, iFaceDiscover: iFaceDiscover, - recorder: peer.NewRecorder(""), + recorder: recorder, ctxCancelLock: &sync.Mutex{}, networkChangeListener: networkChangeListener, + netMgr: netevents.NewManager(recorder), } } @@ -143,16 +192,23 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid defer c.ctxCancel() c.ctxCancelLock.Unlock() - auth := NewAuthWithConfig(ctx, cfg) + auth := NewAuthWithConfig(ctx, cfg, cfgFile) err = auth.login(urlOpener, isAndroidTV) if err != nil { return err } - // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) - connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) - c.setState(cfg, cacheDir, connectClient) + + connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, + internal.WithNetEvents(c.netMgr)) + c.setState(cfg, cacheDir, cfgFile, connectClient) + // 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() + // still reads the previous run's context state, which holds the NeedsLogin + // that prompted this login, and would re-latch what was just cleared. + c.clearLoginRequired() return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } @@ -186,8 +242,9 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) - connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) - c.setState(cfg, cacheDir, connectClient) + connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, + internal.WithNetEvents(c.netMgr)) + c.setState(cfg, cacheDir, cfgFile, connectClient) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } @@ -216,9 +273,49 @@ func (c *Client) RenewTun(fd int) error { return e.RenewTun(fd) } +func (c *Client) GetTunSettings() (*TunSettings, error) { + cc := c.getConnectClient() + if cc == nil { + return nil, fmt.Errorf("engine not running") + } + + e := cc.Engine() + if e == nil { + return nil, fmt.Errorf("engine not initialized") + } + + routes, searchDomains := e.TunSettings() + return &TunSettings{ + Routes: strings.Join(routes, ";"), + SearchDomains: strings.Join(searchDomains, ";"), + }, nil +} + +// SetNetworkAvailable feeds OS-reported network availability into the client. +// While unavailable, the internal reconnect loops suspend their attempts and +// the connection listener reports NoNetwork instead of Connecting; when +// availability returns, the loops resume immediately with a fresh backoff. +// Losing the last network also sweeps the registered connections: nothing can +// redial while offline, so the stale sockets would otherwise stay silently +// "connected" until their own timeouts and the client would keep reporting +// Connected with no network at all. +func (c *Client) SetNetworkAvailable(available bool) { + c.netMgr.SetNetworkAvailable(available) +} + +// NotifyNetworkChange marks the management, signal and relay connections +// stale after the OS switched networks and schedules a sweep that cuts +// whatever has not redialed on the new network by then. The engine and the +// TUN device stay untouched. +func (c *Client) NotifyNetworkChange() { + c.netMgr.NotifyNetworkChange() +} + // DebugBundle generates a debug bundle, uploads it, and returns the upload key. -// It works both with and without a running engine. -func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (string, error) { +// It works both with and without a running engine. anonymizeLevel is "default" +// 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) { cfg, cacheDir, cc := c.stateSnapshot() // If the engine hasn't been started, load config from disk @@ -237,6 +334,7 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin InternalConfig: cfg, StatusRecorder: c.recorder, TempDir: cacheDir, + StatePath: platformFiles.StateFilePath(), } if cc != nil { @@ -260,6 +358,7 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin deps, debug.BundleConfig{ Anonymize: anonymize, + AnonymizeLevel: nbAnonymize.ParseLevel(anonymizeLevel), IncludeSystemInfo: true, }, ) @@ -277,7 +376,7 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() - key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path) + key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path, false) if err != nil { return "", fmt.Errorf("upload debug bundle: %w", err) } @@ -299,6 +398,13 @@ func (c *Client) SetInfoLogLevel() { // PeersList return with the list of the PeerInfos func (c *Client) PeersList() *PeerInfoArray { + // The recorder only caches transfer counters and handshake times; nothing + // refreshes them on its own, so without this they read as zero. The desktop + // daemon does the same before serving a full peer status. + if err := c.recorder.RefreshWireGuardStats(); err != nil { + log.Debugf("failed to refresh WireGuard stats: %v", err) + } + fullStatus := c.recorder.GetFullStatus() peerInfos := make([]PeerInfo, len(fullStatus.Peers)) @@ -309,6 +415,20 @@ func (c *Client) PeersList() *PeerInfoArray { FQDN: p.FQDN, ConnStatus: int(p.ConnStatus), Routes: PeerRoutes{routes: maps.Keys(p.GetRoutes())}, + + PubKey: p.PubKey, + Latency: formatDuration(p.Latency), + LatencyMs: p.Latency.Milliseconds(), + BytesRx: p.BytesRx, + BytesTx: p.BytesTx, + ConnStatusUpdate: formatTime(p.ConnStatusUpdate), + Relayed: p.Relayed, + RosenpassEnabled: p.RosenpassEnabled, + LastWireguardHandshake: formatTime(p.LastWireguardHandshake), + LocalIceCandidateType: p.LocalIceCandidateType, + RemoteIceCandidateType: p.RemoteIceCandidateType, + LocalIceCandidateEndpoint: p.LocalIceCandidateEndpoint, + RemoteIceCandidateEndpoint: p.RemoteIceCandidateEndpoint, } peerInfos[n] = pi } @@ -431,7 +551,11 @@ func (c *Client) OnUpdatedHostDNS(list *DNSList) error { // SetConnectionListener set the network connection listener func (c *Client) SetConnectionListener(listener ConnectionListener) { - c.recorder.SetConnectionListener(listener) + if listener == nil { + c.recorder.RemoveConnectionListener() + return + } + c.recorder.SetConnectionListener(connectionListenerAdapter{listener}) } // RemoveConnectionListener remove connection listener @@ -439,10 +563,6 @@ func (c *Client) RemoveConnectionListener() { c.recorder.RemoveConnectionListener() } -func (c *Client) toggleRoute(command routeCommand) error { - return command.toggleRoute() -} - func (c *Client) getRouteManager() (routemanager.Manager, error) { client := c.getConnectClient() if client == nil { @@ -462,22 +582,22 @@ func (c *Client) getRouteManager() (routemanager.Manager, error) { return manager, nil } -func (c *Client) SelectRoute(route string) error { +func (c *Client) SelectRoute(id string) error { manager, err := c.getRouteManager() if err != nil { return err } - return c.toggleRoute(selectRouteCommand{route: route, manager: manager}) + return manager.SelectRoutes([]route.NetID{route.NetID(id)}, true) } -func (c *Client) DeselectRoute(route string) error { +func (c *Client) DeselectRoute(id string) error { manager, err := c.getRouteManager() if err != nil { return err } - return c.toggleRoute(deselectRouteCommand{route: route, manager: manager}) + return manager.DeselectRoutes([]route.NetID{route.NetID(id)}) } // getNetworkDomainsFromRoute extracts domains from a route and enriches each domain @@ -512,3 +632,28 @@ func exportEnvList(list *EnvList) { } } } + +// formatDuration renders a duration for display, trimming the fractional part +// to two digits so latencies read as "12.34ms" rather than "12.345678ms". +func formatDuration(d time.Duration) string { + ds := d.String() + dotIndex := strings.Index(ds, ".") + if dotIndex == -1 { + return ds + } + + endIndex := min(dotIndex+3, len(ds)) + + // Skip the remaining digits so only the unit suffix is appended back. + unitStart := endIndex + for unitStart < len(ds) && ds[unitStart] >= '0' && ds[unitStart] <= '9' { + unitStart++ + } + return ds[:endIndex] + ds[unitStart:] +} + +// formatTime renders a timestamp in UTC using a fixed layout. The zero time is +// passed through as-is so the UI can recognise it and show "never" instead. +func formatTime(t time.Time) string { + return t.UTC().Format("2006-01-02 15:04:05") +} diff --git a/client/android/connection_listener.go b/client/android/connection_listener.go new file mode 100644 index 000000000..77c47574b --- /dev/null +++ b/client/android/connection_listener.go @@ -0,0 +1,41 @@ +//go:build android + +package android + +import ( + "github.com/netbirdio/netbird/client/internal/peer" +) + +// Client state values delivered via ConnectionListener.OnStateChanged, +// re-exported as basic constants so gomobile emits them into the generated +// Java bindings. They mirror peer.ClientState*: append-only, never reorder. +const ( + ClientStateDisconnected = int(peer.ClientStateDisconnected) + ClientStateConnected = int(peer.ClientStateConnected) + ClientStateConnecting = int(peer.ClientStateConnecting) + ClientStateDisconnecting = int(peer.ClientStateDisconnecting) + ClientStateNoNetwork = int(peer.ClientStateNoNetwork) +) + +// ConnectionListener export internal Listener for mobile. It mirrors +// peer.Listener with OnStateChanged taking a plain int (one of the +// ClientState* constants), because gomobile cannot bind named types. +type ConnectionListener interface { + OnStateChanged(state int) + OnConnected() + OnDisconnected() + OnConnecting() + OnDisconnecting() + OnAddressChanged(string, string) + OnPeersListChanged(int) +} + +// connectionListenerAdapter adapts the gomobile-facing ConnectionListener to +// peer.Listener, converting the typed state to the int the binding carries. +type connectionListenerAdapter struct { + ConnectionListener +} + +func (a connectionListenerAdapter) OnStateChanged(state peer.ClientState) { + a.ConnectionListener.OnStateChanged(int(state)) +} diff --git a/client/android/login.go b/client/android/login.go index a9422cdbf..24c911eb5 100644 --- a/client/android/login.go +++ b/client/android/login.go @@ -4,6 +4,8 @@ import ( "context" "fmt" + log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/system" @@ -36,12 +38,20 @@ type Auth struct { } // NewAuth instantiate Auth struct and validate the management URL +// +// The configuration at cfgPath is reused when one is already there, and only created when it is +// not. Building a fresh in-memory config unconditionally gives the client a new WireGuard key on +// every call: the peer registers under that key, the key is written out, and any peer registered by +// an earlier call is orphaned on the server. It also breaks a client that enrols and then runs from +// the persisted config, because the identity it registered is not the one it runs with — the +// management stream rejects it with "no peer auth method provided". func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { inputCfg := profilemanager.ConfigInput{ + ConfigPath: cfgPath, ManagementURL: mgmURL, } - cfg, err := profilemanager.CreateInMemoryConfig(inputCfg) + cfg, err := profilemanager.UpdateOrCreateConfig(inputCfg) if err != nil { return nil, err } @@ -53,11 +63,14 @@ func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { }, nil } -// NewAuthWithConfig instantiate Auth based on existing config -func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config) *Auth { +// NewAuthWithConfig instantiate Auth based on existing config. cfgPath is the +// file the config was loaded from; it identifies the profile whose account email +// backs the login_hint. +func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config, cfgPath string) *Auth { return &Auth{ - ctx: ctx, - config: config, + ctx: ctx, + config: config, + cfgPath: cfgPath, } } @@ -150,12 +163,14 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error { } jwtToken := "" + email := "" if needsLogin { tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV) if err != nil { return fmt.Errorf("interactive sso login failed: %v", err) } jwtToken = tokenInfo.GetTokenToUse() + email = tokenInfo.Email } err, _ = authClient.Login(a.ctx, "", jwtToken) @@ -163,27 +178,62 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error { return fmt.Errorf("login failed: %v", err) } + // Stored after Login, not before: a rejected token must not leave a hint + // pointing at an account that cannot be used. + if email != "" && a.cfgPath != "" { + if err := writeProfileEmail(a.cfgPath, email); err != nil { + log.Warnf("failed to store profile account email: %v", err) + } + } + go urlOpener.OnLoginSuccess() return nil } func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) { - oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV) + oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV, profileLoginHint(a.cfgPath)) if err != nil { return nil, fmt.Errorf("failed to get OAuth flow: %v", err) } - flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO()) + return runOAuthFlow(a.ctx, oAuthFlow, urlOpener, nil) +} + +// profileLoginHint returns the stored account email for the profile at cfgPath. +// An empty hint is deliberate, not a fallback: a fresh profile leaves the +// choice to the IdP. Switching accounts is done by switching or removing +// profiles, not by logging out — logout keeps the email. +func profileLoginHint(cfgPath string) string { + if cfgPath == "" { + return "" + } + return readProfileEmail(cfgPath) +} + +// runOAuthFlow drives an already acquired OAuth flow to a token: requests the +// flow info, presents the verification URL through the opener and waits for +// the browser round-trip. Open is called synchronously — it is what marks the +// surface as opened on the client side, and a fast token's OnLoginSuccess is +// a no-op until it has, so the dismissal would be dropped rather than +// delayed. Openers must therefore not block: they post their UI work and +// return. onWaiting, when set, runs after the URL is shown, right before the +// blocking wait. +func runOAuthFlow(ctx context.Context, flow auth.OAuthFlow, urlOpener URLOpener, onWaiting func()) (*auth.TokenInfo, error) { + flowInfo, err := flow.RequestAuthInfo(ctx) if err != nil { - return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err) + return nil, fmt.Errorf("request auth info: %w", err) } - go urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode) + urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode) - tokenInfo, err := oAuthFlow.WaitToken(a.ctx, flowInfo) + if onWaiting != nil { + onWaiting() + } + + tokenInfo, err := flow.WaitToken(ctx, flowInfo) if err != nil { - return nil, fmt.Errorf("waiting for browser login failed: %v", err) + return nil, fmt.Errorf("wait for token: %w", err) } return &tokenInfo, nil diff --git a/client/android/login_test.go b/client/android/login_test.go new file mode 100644 index 000000000..b04790f6b --- /dev/null +++ b/client/android/login_test.go @@ -0,0 +1,51 @@ +package android + +import ( + "path/filepath" + "testing" +) + +// NewAuth must reuse the configuration already at cfgPath rather than building a fresh one. +// +// Creating a new in-memory config on every call gives the client a new WireGuard private key each +// time. The peer registers under that key and the key is written out, so a peer registered by an +// earlier call is orphaned on the server — a client that enrols twice leaves two entries and owns +// neither. It also breaks enrol-then-run: RunWithoutLogin reloads the configuration from disk, so +// the identity that registered is not the identity that runs, and the management stream rejects it +// with "no peer auth method provided, please use a setup key or interactive SSO login". +func TestNewAuth_ReusesPersistedIdentity(t *testing.T) { + cfgPath := filepath.Join(t.TempDir(), "config.json") + + first, err := NewAuth(cfgPath, "https://api.example.com:443") + if err != nil { + t.Fatalf("first NewAuth: %v", err) + } + if first.config.PrivateKey == "" { + t.Fatal("first NewAuth produced no private key") + } + + second, err := NewAuth(cfgPath, "https://api.example.com:443") + if err != nil { + t.Fatalf("second NewAuth: %v", err) + } + + if second.config.PrivateKey != first.config.PrivateKey { + t.Errorf("private key changed between calls: a second enrolment would orphan the peer registered by the first") + } +} + +// A missing configuration is still created, so a first enrolment works unchanged. +func TestNewAuth_CreatesConfigWhenAbsent(t *testing.T) { + cfgPath := filepath.Join(t.TempDir(), "config.json") + + auth, err := NewAuth(cfgPath, "https://api.example.com:443") + if err != nil { + t.Fatalf("NewAuth: %v", err) + } + if auth.config == nil || auth.config.PrivateKey == "" { + t.Fatal("NewAuth did not create a usable configuration") + } + if auth.cfgPath != cfgPath { + t.Errorf("cfgPath = %q, want %q", auth.cfgPath, cfgPath) + } +} diff --git a/client/android/peer_notifier.go b/client/android/peer_notifier.go index c2595e574..f525055bb 100644 --- a/client/android/peer_notifier.go +++ b/client/android/peer_notifier.go @@ -12,12 +12,30 @@ const ( ) // PeerInfo describe information about the peers. It designed for the UI usage +// +// The fields below ConnStatus back the peer detail screen. Durations and times +// are pre-formatted into strings so the UI does not have to know Go's layouts; +// Latency is additionally exposed as LatencyMs for colour coding. type PeerInfo struct { IP string IPv6 string FQDN string ConnStatus int Routes PeerRoutes + + PubKey string + Latency string + LatencyMs int64 + BytesRx int64 + BytesTx int64 + ConnStatusUpdate string + Relayed bool + RosenpassEnabled bool + LastWireguardHandshake string + LocalIceCandidateType string + RemoteIceCandidateType string + LocalIceCandidateEndpoint string + RemoteIceCandidateEndpoint string } func (p *PeerInfo) GetPeerRoutes() *PeerRoutes { diff --git a/client/android/profile_manager.go b/client/android/profile_manager.go index 87c001396..20d585d6a 100644 --- a/client/android/profile_manager.go +++ b/client/android/profile_manager.go @@ -13,18 +13,18 @@ import ( ) const ( - // Android-specific config filename (different from desktop default.json) - defaultConfigFilename = "netbird.cfg" - // Subdirectory for non-default profiles (must match Java Preferences.java) - profilesSubdir = "profiles" // Android uses a single user context per app (non-empty username required by ServiceManager) androidUsername = "android" ) // Profile represents a profile for gomobile type Profile struct { - ID string - Name string + ID string + Name string + // Email is the account this profile last logged in with, "" if it never + // completed an SSO login. Kept across logouts; cleared when the profile is + // removed. See profile_state.go. + Email string IsActive bool } @@ -101,6 +101,7 @@ func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) { profiles = append(profiles, &Profile{ ID: p.ID.String(), Name: p.Name, + Email: pm.profileEmail(p.ID.String()), IsActive: p.IsActive, }) } @@ -123,7 +124,22 @@ func (pm *ProfileManager) GetActiveProfile() (*Profile, error) { if err != nil { return nil, fmt.Errorf("failed to resolve active profile %q: %w", activeState.ID, err) } - return &Profile{ID: prof.ID.String(), Name: prof.Name, IsActive: true}, nil + return &Profile{ + ID: prof.ID.String(), + Name: prof.Name, + Email: pm.profileEmail(prof.ID.String()), + IsActive: true, + }, nil +} + +// profileEmail returns the account email recorded for a profile. Display-only, so +// an unresolvable path degrades to "" rather than an error. +func (pm *ProfileManager) profileEmail(id string) string { + configPath, err := pm.getProfileConfigPath(id) + if err != nil { + return "" + } + return readProfileEmail(configPath) } // SwitchProfile switches to a different profile @@ -185,17 +201,46 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return fmt.Errorf("failed to save config: %w", err) } + // The stored account email is kept on purpose, matching the desktop and CLI + // logout semantics: the next login passes it as the login_hint so the IdP + // preselects the account. Removing the profile is what deletes it. log.Infof("logged out from profile: %s", id) return nil } +// RenameProfile changes a profile's display name. The profile ID, and therefore +// its on-disk filename, is left untouched: only the "name" field of the config +// is rewritten. This works for the default profile too, whose config lives in +// netbird.cfg rather than under profiles/. +func (pm *ProfileManager) RenameProfile(id string, newName string) error { + if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), androidUsername, newName); err != nil { + return fmt.Errorf("failed to rename profile: %w", err) + } + + log.Infof("renamed profile %s to: %s", id, newName) + return nil +} + // RemoveProfile deletes a profile func (pm *ProfileManager) RemoveProfile(id string) error { + configPath, err := pm.getProfileConfigPath(id) + if err != nil { + return err + } + // Use ServiceManager (removes profile from profiles/ directory) if err := pm.serviceMgr.RemoveProfile(profilemanager.ID(id), androidUsername); err != nil { return fmt.Errorf("failed to remove profile: %w", err) } + // The account file is this package's, not the ServiceManager's, so it must + // go here. The default profile has a fixed filename, so a recreated one + // would otherwise inherit the deleted profile's email as its login_hint. + // Not fatal: the profile itself is gone. + if err := removeProfileEmail(configPath); err != nil { + log.Warnf("failed to remove stored account email for profile %s: %v", id, err) + } + log.Infof("removed profile: %s", id) return nil } diff --git a/client/android/profile_prefs.go b/client/android/profile_prefs.go new file mode 100644 index 000000000..9c1fd307b --- /dev/null +++ b/client/android/profile_prefs.go @@ -0,0 +1,38 @@ +//go:build android + +package android + +import ( + "fmt" + + "github.com/netbirdio/netbird/client/internal/profilemanager" +) + +type prefsStore interface { + Get(namespace string, v any) (bool, error) + Put(namespace string, v any) error +} + +type profilePrefs struct { + prefs *profilemanager.Prefs +} + +func newProfilePrefs(configDir, profileID string) (*profilePrefs, error) { + if configDir == "" || profileID == "" { + return nil, fmt.Errorf("profile prefs require a config dir and profile ID") + } + pm := NewProfileManager(configDir) + prefs, err := pm.serviceMgr.ProfilePrefs(profilemanager.ID(profileID), androidUsername) + if err != nil { + return nil, fmt.Errorf("resolve profile prefs: %w", err) + } + return &profilePrefs{prefs: prefs}, nil +} + +func (p *profilePrefs) Get(namespace string, v any) (bool, error) { + return p.prefs.Get(namespace, v) +} + +func (p *profilePrefs) Put(namespace string, v any) error { + return p.prefs.Put(namespace, v) +} diff --git a/client/android/profile_state.go b/client/android/profile_state.go new file mode 100644 index 000000000..0063b587f --- /dev/null +++ b/client/android/profile_state.go @@ -0,0 +1,108 @@ +package android + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/util" +) + +const ( + // Android-specific config filename (different from desktop default.json) + defaultConfigFilename = "netbird.cfg" + // Subdirectory for non-default profiles (must match Java Preferences.java) + profilesSubdir = "profiles" + // profileAccountSuffix names the file holding the profile's account email. + // Deliberately not ".state.json", which desktop uses for the same data: + // there the email and the engine's state manager live in different + // directories, but on Android both resolve under files/, so sharing the name + // would have the two overwrite each other — the state manager rewrites the + // whole file from its own keys (see statemanager.Manager.PersistState), and + // this package's writer does the same in reverse. + profileAccountSuffix = ".account.json" +) + +// profileAccountPathFor derives the account file path from a profile's config +// path: netbird.cfg -> netbird.account.json, .json -> .account.json. +// +// Deriving from the config path rather than resolving the active profile keeps +// the write on the profile the login actually ran for: Auth.login runs in a +// goroutine, so the active profile can change under a flow already in flight. +func profileAccountPathFor(configPath string) (string, error) { + if configPath == "" { + return "", fmt.Errorf("empty config path") + } + + base := filepath.Base(configPath) + stem := strings.TrimSuffix(base, filepath.Ext(base)) + if stem == "" || stem == "." { + return "", fmt.Errorf("config path %q has no filename stem", configPath) + } + + return filepath.Join(filepath.Dir(configPath), stem+profileAccountSuffix), nil +} + +// readProfileEmail returns the account email stored for the profile whose config +// lives at configPath. A missing or unreadable file yields "", which leaves the +// account choice to the IdP. +func readProfileEmail(configPath string) string { + accountPath, err := profileAccountPathFor(configPath) + if err != nil { + log.Debugf("no profile account path for login hint: %v", err) + return "" + } + + var state profilemanager.ProfileState + if _, err := util.ReadJson(accountPath, &state); err != nil { + if !os.IsNotExist(err) { + log.Debugf("failed to read profile account for login hint: %v", err) + } + return "" + } + + return state.Email +} + +// writeProfileEmail records the account email for the profile whose config lives +// at configPath, so later logins can pass it as an OIDC login_hint. An empty +// email is ignored rather than blanking what is already stored. +func writeProfileEmail(configPath string, email string) error { + if email == "" { + return nil + } + + accountPath, err := profileAccountPathFor(configPath) + if err != nil { + return fmt.Errorf("resolve profile account path: %w", err) + } + + state := profilemanager.ProfileState{Email: email} + if err := util.WriteJsonWithRestrictedPermission(context.Background(), accountPath, state); err != nil { + return fmt.Errorf("write profile account: %w", err) + } + + return nil +} + +// removeProfileEmail drops the stored account email. Called on profile removal, +// not on logout: a logged-out profile keeps its email so the next login passes +// it as the login_hint, matching the desktop and CLI semantics. Mirrors the +// desktop UI's RemoveProfileState call. +func removeProfileEmail(configPath string) error { + accountPath, err := profileAccountPathFor(configPath) + if err != nil { + return fmt.Errorf("resolve profile account path: %w", err) + } + + if err := os.Remove(accountPath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("remove profile account: %w", err) + } + + return nil +} diff --git a/client/android/profile_state_test.go b/client/android/profile_state_test.go new file mode 100644 index 000000000..82a1c2a87 --- /dev/null +++ b/client/android/profile_state_test.go @@ -0,0 +1,161 @@ +package android + +import ( + "os" + "path/filepath" + "testing" +) + +func TestProfileAccountPathFor(t *testing.T) { + tests := []struct { + name string + configPath string + want string + wantErr bool + }{ + { + name: "default profile", + configPath: "/data/data/io.netbird.client/files/netbird.cfg", + want: filepath.FromSlash("/data/data/io.netbird.client/files/netbird.account.json"), + }, + { + name: "id profile", + configPath: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.json", + want: filepath.FromSlash("/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.account.json"), + }, + { + name: "legacy name-keyed profile is handled the same way", + configPath: "/data/data/io.netbird.client/files/profiles/work.json", + want: filepath.FromSlash("/data/data/io.netbird.client/files/profiles/work.account.json"), + }, + { + name: "empty path is rejected", + configPath: "", + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := profileAccountPathFor(tt.configPath) + if tt.wantErr { + if err == nil { + t.Fatalf("expected an error, got path %q", got) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != tt.want { + t.Errorf("got %q, want %q", got, tt.want) + } + }) + } +} + +func TestProfileAccountPathForDefaultDoesNotCollide(t *testing.T) { + root := "/data/data/io.netbird.client/files" + + defaultAccount, err := profileAccountPathFor(filepath.Join(root, defaultConfigFilename)) + if err != nil { + t.Fatalf("default profile: %v", err) + } + + idAccount, err := profileAccountPathFor(filepath.Join(root, profilesSubdir, "abc123.json")) + if err != nil { + t.Fatalf("id profile: %v", err) + } + + if defaultAccount == idAccount { + t.Fatalf("default and id profile share an account file: %q", defaultAccount) + } +} + +// The account file must never land on the engine state file: on Android both +// resolve under files/, and the state manager rewrites the whole file from its +// own keys, so sharing a path would have the two overwrite each other. The +// expected names here mirror ProfileManager.GetStateFilePath. +func TestProfileAccountPathAvoidsEngineStateFile(t *testing.T) { + root := "/data/data/io.netbird.client/files" + + cases := []struct { + configPath string + engineState string + }{ + { + configPath: filepath.Join(root, defaultConfigFilename), + engineState: filepath.Join(root, "state.json"), + }, + { + configPath: filepath.Join(root, profilesSubdir, "abc123.json"), + engineState: filepath.Join(root, profilesSubdir, "abc123.state.json"), + }, + } + + for _, c := range cases { + account, err := profileAccountPathFor(c.configPath) + if err != nil { + t.Fatalf("%s: %v", c.configPath, err) + } + if account == c.engineState { + t.Errorf("account file collides with the engine state file: %q", account) + } + } +} + +func TestWriteThenReadProfileEmail(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json") + if err := ensureDirFor(t, configPath); err != nil { + t.Fatalf("prepare dir: %v", err) + } + + if got := readProfileEmail(configPath); got != "" { + t.Errorf("expected no email before a login, got %q", got) + } + + const email = "user@example.com" + if err := writeProfileEmail(configPath, email); err != nil { + t.Fatalf("write: %v", err) + } + + if got := readProfileEmail(configPath); got != email { + t.Errorf("got %q, want %q", got, email) + } + + if err := removeProfileEmail(configPath); err != nil { + t.Fatalf("remove: %v", err) + } + if got := readProfileEmail(configPath); got != "" { + t.Errorf("expected no email after removal, got %q", got) + } + + // Removal may run on a never-logged-in profile, so a second remove must pass. + if err := removeProfileEmail(configPath); err != nil { + t.Fatalf("second remove should be a no-op: %v", err) + } +} + +func TestWriteProfileEmailIgnoresEmpty(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json") + if err := ensureDirFor(t, configPath); err != nil { + t.Fatalf("prepare dir: %v", err) + } + + const email = "user@example.com" + if err := writeProfileEmail(configPath, email); err != nil { + t.Fatalf("write: %v", err) + } + if err := writeProfileEmail(configPath, ""); err != nil { + t.Fatalf("write empty: %v", err) + } + + if got := readProfileEmail(configPath); got != email { + t.Errorf("empty write clobbered the stored email: got %q, want %q", got, email) + } +} + +func ensureDirFor(t *testing.T, path string) error { + t.Helper() + return os.MkdirAll(filepath.Dir(path), 0o700) +} diff --git a/client/android/route_command.go b/client/android/route_command.go deleted file mode 100644 index 5e7357335..000000000 --- a/client/android/route_command.go +++ /dev/null @@ -1,70 +0,0 @@ -//go:build android - -package android - -import ( - "fmt" - - log "github.com/sirupsen/logrus" - "golang.org/x/exp/maps" - - "github.com/netbirdio/netbird/client/internal/routemanager" - "github.com/netbirdio/netbird/route" -) - -func executeRouteToggle(id string, manager routemanager.Manager, - operationName string, - routeOperation func(routes []route.NetID, allRoutes []route.NetID) error) error { - netID := route.NetID(id) - routes := []route.NetID{netID} - - routesMap := manager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - - log.Debugf("%s with ids: %v", operationName, routes) - - if err := routeOperation(routes, maps.Keys(routesMap)); err != nil { - log.Debugf("error when %s: %s", operationName, err) - return fmt.Errorf("error %s: %w", operationName, err) - } - - manager.TriggerSelection(manager.GetClientRoutes()) - - return nil -} - -type routeCommand interface { - toggleRoute() error -} - -type selectRouteCommand struct { - route string - manager routemanager.Manager -} - -func (s selectRouteCommand) toggleRoute() error { - routeSelector := s.manager.GetRouteSelector() - if routeSelector == nil { - return fmt.Errorf("no route selector available") - } - - routeOperation := func(routes []route.NetID, allRoutes []route.NetID) error { - return routeSelector.SelectRoutes(routes, true, allRoutes) - } - - return executeRouteToggle(s.route, s.manager, "selecting route", routeOperation) -} - -type deselectRouteCommand struct { - route string - manager routemanager.Manager -} - -func (d deselectRouteCommand) toggleRoute() error { - routeSelector := d.manager.GetRouteSelector() - if routeSelector == nil { - return fmt.Errorf("no route selector available") - } - - return executeRouteToggle(d.route, d.manager, "deselecting route", routeSelector.DeselectRoutes) -} diff --git a/client/android/session.go b/client/android/session.go new file mode 100644 index 000000000..d5da09c93 --- /dev/null +++ b/client/android/session.go @@ -0,0 +1,312 @@ +//go:build android + +package android + +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. +// +// OnStateChanged is a payload-free wake-up whenever the state snapshot +// 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. +type StateChangeListener interface { + OnStateChanged() + OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool) +} + +// Status returns the connect run-loop's status label — the same value the +// desktop daemon serves in StatusResponse.Status. "NeedsLogin" means the +// management server rejected the peer and an interactive login is required. +// +// The label is latched: the run loop keeps its status in a per-run context +// state, which a restart replaces with a fresh Idle one, so an engine restart +// (network change, always-on) would otherwise erase the fact that the peer +// still needs to log in. Only a successful interactive login or extend clears +// it — see clearLoginRequired. +func (c *Client) Status() string { + latched, generation := c.loginRequiredState() + if latched { + return string(internal.StatusNeedsLogin) + } + cc := c.getConnectClient() + if cc == nil { + return string(internal.StatusIdle) + } + status := cc.Status() + if status == internal.StatusNeedsLogin { + c.latchLoginRequired(generation) + } + return string(status) +} + +func (c *Client) loginRequiredState() (bool, uint64) { + c.loginRequiredMu.Lock() + defer c.loginRequiredMu.Unlock() + return c.loginRequired, c.loginCleared +} + +// latchLoginRequired records a NeedsLogin observation, unless a clear landed +// while the caller was reading the run loop's status: cc.Status() is read +// outside the lock, so a login or extend completing in that window would +// otherwise be undone by this stale observation, stranding the UI on +// "login required" over a healthy session. +func (c *Client) latchLoginRequired(observedGeneration uint64) { + c.loginRequiredMu.Lock() + defer c.loginRequiredMu.Unlock() + if c.loginCleared != observedGeneration { + return + } + c.loginRequired = true +} + +// clearLoginRequired releases the latch after a successful interactive login +// or session extend, and invalidates any observation already in flight. +func (c *Client) clearLoginRequired() { + c.loginRequiredMu.Lock() + defer c.loginRequiredMu.Unlock() + c.loginRequired = false + c.loginCleared++ +} + +// SessionExpiresAtUnix returns the SSO session deadline as unix seconds, or 0 +// when no deadline is known (not SSO-registered, expiry disabled, or the +// engine has not received one yet). A past value means the session expired. +// Mirror of StatusResponse.sessionExpiresAt on the desktop daemon. +func (c *Client) SessionExpiresAtUnix() int64 { + deadline := c.recorder.GetSessionExpiresAt() + if deadline.IsZero() { + return 0 + } + return deadline.Unix() +} + +// SetStateChangeListener registers the state notification listener. +// Replaces any previously registered listener; remove it with +// RemoveStateChangeListener. +func (c *Client) SetStateChangeListener(listener StateChangeListener) { + c.stateChangeMu.Lock() + defer c.stateChangeMu.Unlock() + c.stopStateChangeWatchLocked() + if listener == nil { + 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. + done := make(chan struct{}) + c.stateChangeDone = done + + id, ch := c.recorder.SubscribeToStateChanges() + c.stateChangeSubID = id + // The channel is closed by UnsubscribeFromStateChanges, which ends the + // goroutine. Ticks are coalesced (buffer of one), so a burst of changes + // wakes the listener once. + go func() { + for range ch { + select { + case <-done: + return + default: + } + listener.OnStateChanged() + } + }() + + c.eventSub = c.recorder.SubscribeToEvents() + go watchSessionWarnings(c.eventSub, listener, done) +} + +// RemoveStateChangeListener unregisters the state notification listener. +func (c *Client) RemoveStateChangeListener() { + c.stateChangeMu.Lock() + defer c.stateChangeMu.Unlock() + 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 +// listener. Mirror of the daemon's RequestExtendAuthSession / +// WaitExtendAuthSession RPC pair, with URLOpener playing the "UI opens the +// browser" role. +// +// Only one flow may be in flight: the PKCE step binds a fixed loopback port, +// so a second concurrent flow would fail on that bind. Call +// CancelExtendAuthSession when the user abandons the browser. +func (c *Client) ExtendAuthSession(urlOpener URLOpener, isAndroidTV bool, resultListener ErrListener) { + ctx, err := c.beginExtend() + if err != nil { + resultListener.OnError(err) + return + } + + go func() { + defer c.endExtend() + if err := c.extendAuthSession(ctx, urlOpener, isAndroidTV); err != nil { + resultListener.OnError(err) + return + } + resultListener.OnSuccess() + }() +} + +// CancelExtendAuthSession aborts an in-flight ExtendAuthSession. The tunnel is +// left alone — unlike the login flow, which cancels the whole client context +// by stopping the engine. Without this the abandoned PKCE wait keeps its +// loopback port for the full flow timeout and blocks every later attempt. +// No-op when no flow is running. +func (c *Client) CancelExtendAuthSession() { + c.extendMu.Lock() + defer c.extendMu.Unlock() + if c.extendCancel != nil { + c.extendCancel() + } +} + +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 + // that is no longer registered. + if c.stateChangeDone != nil { + close(c.stateChangeDone) + c.stateChangeDone = nil + } + if c.stateChangeSubID != "" { + 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) { + c.extendMu.Lock() + defer c.extendMu.Unlock() + if c.extendCancel != nil { + return nil, fmt.Errorf("session extend already in progress") + } + ctx, cancel := context.WithCancel(context.Background()) + c.extendCancel = cancel + return ctx, nil +} + +func (c *Client) endExtend() { + c.extendMu.Lock() + defer c.extendMu.Unlock() + if c.extendCancel != nil { + c.extendCancel() + c.extendCancel = nil + } +} + +func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isAndroidTV bool) error { + cfg, cfgPath, cc := c.authSnapshot() + if cfg == nil || cc == nil { + return fmt.Errorf("engine is not running") + } + engine := cc.Engine() + if engine == nil { + return fmt.Errorf("engine is not initialized") + } + + authClient, err := auth.NewAuth(ctx, cfg.PrivateKey, cfg.ManagementURL, cfg) + if err != nil { + return fmt.Errorf("failed to create auth client: %v", err) + } + defer authClient.Close() + + // Passing the config path makes the flow pick up the login_hint: an extend + // renews the session of the account already signed in, so it must not stop to + // offer a choice. + a := NewAuthWithConfig(ctx, cfg, cfgPath) + tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV) + if err != nil { + return fmt.Errorf("interactive sso login failed: %v", err) + } + + if _, err := engine.ExtendAuthSession(ctx, tokenInfo.GetTokenToUse()); err != nil { + return err + } + c.clearLoginRequired() + + go urlOpener.OnLoginSuccess() + return nil +} diff --git a/client/android/ssh_client.go b/client/android/ssh_client.go new file mode 100644 index 000000000..2822b6539 --- /dev/null +++ b/client/android/ssh_client.go @@ -0,0 +1,649 @@ +//go:build android + +package android + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "strconv" + "strings" + "sync" + "time" + + log "github.com/sirupsen/logrus" + gossh "golang.org/x/crypto/ssh" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/internal/auth" + "github.com/netbirdio/netbird/client/internal/profilemanager" + nbssh "github.com/netbirdio/netbird/client/ssh" + "github.com/netbirdio/netbird/client/ssh/detection" +) + +const ( + sshDialTimeout = 30 * time.Second + sshDetectionTimeout = 5 * time.Second +) + +// 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. +const PasswordRequiredMarker = "netbird-ssh-password-required" + +// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation, +// retry with TrustHostKey set. The presented fingerprint is appended after the +// marker so the prompt can display it and the retry can guard against a key +// that changed between the two connects. Only regular (non-NetBird) servers +// reach this: NetBird peers verify against the registry. +const HostKeyUnknownMarker = "netbird-ssh-hostkey-unknown" + +var ( + errPasswordRequired = errors.New(PasswordRequiredMarker) + errClientClosed = errors.New("ssh client closed") +) + +// errHostKeyUnknown carries the presented fingerprint so Connect can build the +// marker message the Java side parses. +type errHostKeyUnknown struct { + fingerprint string +} + +func (e *errHostKeyUnknown) Error() string { + return HostKeyUnknownMarker + ":" + e.fingerprint +} + +// SSHTerminalListener receives SSH session events. It is implemented in Java. +// +// All callbacks are invoked from goroutines and may run concurrently with each +// other; the implementation must be safe to call from any thread. +type SSHTerminalListener interface { + OnConnected() + OnData(data []byte) + OnClose(reason string) + OnError(message string) +} + +// SSHClient is a NetBird-aware SSH client exposed to Java via gomobile. +// +// It dials through the running NetBird tunnel and runs a standard SSH session +// on top with PTY enabled. Host-key verification uses the NetBird-provided +// peer SSH host keys, identical to the desktop client. +type SSHClient struct { + nb *Client + mu sync.Mutex + listener SSHTerminalListener + urlOpener URLOpener + + sshClient *gossh.Client + session *gossh.Session + stdin io.WriteCloser + closed bool + + // gen identifies the current connection attempt. Connect and Close bump it, + // so an in-flight dial or a reader left over from a previous connection + // finds itself stale and stays silent instead of publishing OnConnected or + // OnClose for a connection the caller already abandoned. + gen uint64 + dialCancel context.CancelFunc + + // knownHostsConfigDir and knownHostsProfile locate the TOFU store for + // regular SSH servers in the profile's preferences. Java supplies them, + // since an overlay IP is a different host under a different profile. Empty + // until set: without them a regular server cannot be verified and Connect + // refuses one. + knownHostsConfigDir string + knownHostsProfile string + // trustHostKey carries the fingerprint the user confirmed on a previous + // attempt, so the retry accepts exactly that key and persists it. + trustHostKey string +} + +// NewSSHClient creates a new SSH client bound to the running NetBird Client. +func NewSSHClient(c *Client) *SSHClient { + return &SSHClient{nb: c} +} + +// SetListener registers the Java listener. Must be called before Connect to +// receive any events. +func (s *SSHClient) SetListener(l SSHTerminalListener) { + s.mu.Lock() + s.listener = l + s.mu.Unlock() +} + +// SetURLOpener registers the Java URL opener used to display the device-code +// authorization page in a Custom Tabs window when the target peer requires +// JWT authentication. Must be set before Connect to be effective. +func (s *SSHClient) SetURLOpener(opener URLOpener) { + s.mu.Lock() + s.urlOpener = opener + s.mu.Unlock() +} + +// SetKnownHostsStore points the TOFU host-key store at a profile's preferences. +// Must be set before connecting to a regular SSH server; without it such a +// server cannot be verified and Connect refuses one. +func (s *SSHClient) SetKnownHostsStore(configDir, profileID string) { + s.mu.Lock() + s.knownHostsConfigDir = configDir + s.knownHostsProfile = profileID + s.mu.Unlock() +} + +// TrustHostKey records the fingerprint the user confirmed for a regular server, +// so the next Connect accepts that exact key and adds it to the known-hosts +// store. Passing a fingerprint that no longer matches makes the connect fail +// rather than trust a key that changed since the prompt. +func (s *SSHClient) TrustHostKey(fingerprint string) { + s.mu.Lock() + s.trustHostKey = fingerprint + s.mu.Unlock() +} + +// Connect dials the SSH server through the NetBird tunnel and performs the +// SSH handshake. It auto-detects the server type via SSH banner inspection +// and selects the appropriate authentication path: +// +// - NetBird-SSH server requiring JWT: launches the OAuth 2.0 device-code +// flow, opens the verification URL through the registered URLOpener, and +// uses the resulting token as the SSH password. Host-key verification +// uses the NetBird peer registry. +// - NetBird-SSH server without JWT: authenticates with the NetBird SSH +// private key. Host-key verification uses the NetBird peer registry. +// - Regular SSH server (e.g. OpenSSH): authenticates with the NetBird key +// first (so a user-installed NetBird public key works), then falls back +// to the supplied password if non-empty. Host-key verification is +// trust-on-first-use against the per-profile known-hosts store. +// +// The password parameter is only consulted for regular SSH servers. +func (s *SSHClient) Connect(host string, port int, user, password string) error { + if port < 1 || port > 65535 { + return fmt.Errorf("invalid port: %d", port) + } + + cfg, cfgPath, cc := s.nb.authSnapshot() + if cc == nil { + return errors.New("netbird client not running") + } + if cfg == nil { + return errors.New("netbird config not loaded") + } + engine := cc.Engine() + if engine == nil { + return errors.New("netbird engine not available") + } + + s.mu.Lock() + s.gen++ + gen := s.gen + s.mu.Unlock() + + serverType := detectServerType(host, port) + log.Debugf("SSH server type: %s", serverType) + + authMethods, hostKeyCallback, err := s.buildAuth(cfg, cfgPath, engine, serverType, password) + if err != nil { + return err + } + + clientConfig := &gossh.ClientConfig{ + User: user, + Auth: authMethods, + HostKeyCallback: hostKeyCallback, + Timeout: sshDialTimeout, + } + err = s.dialAndHandshake(gen, host, port, clientConfig) + + // An unknown host key is a prompt, not a failure: return the marker intact + // (rootCause would unwrap it) so Java can show the fingerprint and retry. + var unknownHost *errHostKeyUnknown + if errors.As(err, &unknownHost) { + return errors.New(unknownHost.Error()) + } + + // A regular server may still accept a password, so let the caller ask for + // one instead of failing. NetBird servers never use a password, so a + // failure there is genuine. + if err != nil && serverType != detection.ServerTypeNetBirdJWT && + serverType != detection.ServerTypeNetBirdNoJWT && isAuthFailure(err) && + passwordCouldHelp(err, password != "") { + return errPasswordRequired + } + if err != nil { + return rootCause(err) + } + return nil +} + +// StartSession requests a PTY and starts an interactive shell. Output from +// the session is forwarded to the listener via OnData. +func (s *SSHClient) StartSession(cols, rows int) error { + err := s.startSession(cols, rows) + if err != nil { + log.Infof("SSH: start session failed: %v", err) + return rootCause(err) + } + return nil +} + +// Write sends data to the SSH session stdin. +func (s *SSHClient) Write(data []byte) error { + s.mu.Lock() + stdin := s.stdin + s.mu.Unlock() + if stdin == nil { + return errors.New("ssh session not started") + } + if _, err := stdin.Write(data); err != nil { + return fmt.Errorf("write stdin: %w", err) + } + return nil +} + +// Resize updates the PTY window size. +func (s *SSHClient) Resize(cols, rows int) error { + s.mu.Lock() + session := s.session + s.mu.Unlock() + if session == nil { + return errors.New("ssh session not started") + } + return session.WindowChange(rows, cols) +} + +// Reset makes a closed client usable for another Connect: Close leaves the +// one-shot guard set, and clearing it lets the same client back a reconnect. +func (s *SSHClient) Reset() { + s.mu.Lock() + defer s.mu.Unlock() + s.closed = false +} + +// Close terminates the SSH session and underlying connection. Safe to call +// multiple times. +func (s *SSHClient) Close() error { + s.mu.Lock() + s.gen++ + if s.dialCancel != nil { + s.dialCancel() + s.dialCancel = nil + } + sshClient := s.sshClient + session := s.session + stdin := s.stdin + s.sshClient = nil + s.session = nil + s.stdin = nil + notify := !s.closed + s.closed = true + listener := s.listener + s.mu.Unlock() + + if stdin != nil { + if err := stdin.Close(); err != nil { + log.Debugf("ssh: stdin close: %v", err) + } + } + if session != nil { + if err := session.Close(); err != nil && !errors.Is(err, io.EOF) { + log.Debugf("ssh: session close: %v", err) + } + } + var firstErr error + if sshClient != nil { + if err := sshClient.Close(); err != nil { + firstErr = err + } + } + if notify && listener != nil { + listener.OnClose("closed by client") + } + return firstErr +} + +func (s *SSHClient) startSession(cols, rows int) error { + log.Debugf("SSH: starting session %dx%d", cols, rows) + s.mu.Lock() + sshClient := s.sshClient + gen := s.gen + s.mu.Unlock() + + if sshClient == nil { + return errors.New("ssh client not connected") + } + + pty, err := nbssh.StartPTYSession(sshClient, cols, rows) + if err != nil { + return err + } + + s.mu.Lock() + if gen != s.gen { + s.mu.Unlock() + closeQuiet(pty.Session, "stale session") + return errClientClosed + } + s.session = pty.Session + s.stdin = pty.Stdin + s.mu.Unlock() + + readerDone := make(chan string, 2) + go func() { readerDone <- s.readLoop(pty.Stdout, "stdout") }() + go func() { readerDone <- s.readLoop(pty.Stderr, "stderr") }() + go func() { + reason := <-readerDone + if second := <-readerDone; reason == "" { + reason = second + } + s.notifyClose(gen, reason) + }() + log.Debug("SSH: session started, shell running") + return nil +} + +func (s *SSHClient) buildAuth(cfg *profilemanager.Config, cfgPath string, engine *internal.Engine, + serverType detection.ServerType, password string) ([]gossh.AuthMethod, gossh.HostKeyCallback, error) { + + switch serverType { + case detection.ServerTypeNetBirdJWT: + token, err := s.requestJWTToken(cfg, cfgPath) + if err != nil { + return nil, nil, fmt.Errorf("jwt: %w", err) + } + auths := []gossh.AuthMethod{gossh.Password(token)} + return auths, nbssh.CreateHostKeyCallback(nbssh.PeerKeyLookup(engine.GetPeerSSHKey)), nil + + case detection.ServerTypeNetBirdNoJWT: + if cfg.SSHKey == "" { + return nil, nil, errors.New("no NetBird SSH key available") + } + signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey)) + if err != nil { + return nil, nil, fmt.Errorf("parse netbird ssh key: %w", err) + } + auths := []gossh.AuthMethod{gossh.PublicKeys(signer)} + return auths, nbssh.CreateHostKeyCallback(nbssh.PeerKeyLookup(engine.GetPeerSSHKey)), nil + + case detection.ServerTypeRegular: + var auths []gossh.AuthMethod + if cfg.SSHKey != "" { + if signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey)); err == nil { + auths = append(auths, gossh.PublicKeys(signer)) + } else { + log.Debugf("ssh: parse netbird key for regular auth: %v", err) + } + } + if password != "" { + pw := password + auths = append(auths, gossh.Password(pw)) + auths = append(auths, gossh.KeyboardInteractive(func(_, _ string, questions []string, _ []bool) ([]string, error) { + answers := make([]string, len(questions)) + for i := range questions { + answers[i] = pw + } + return answers, nil + })) + } + if len(auths) == 0 { + // Nothing to offer at all: ask for a password rather than failing, + // so the caller can retry once the user supplies one. + return nil, nil, errPasswordRequired + } + callback, err := s.tofuHostKeyCallback() + if err != nil { + return nil, nil, err + } + return auths, callback, nil + + default: + return nil, nil, fmt.Errorf("unsupported SSH server type: %v", serverType) + } +} + +// tofuHostKeyCallback verifies a regular server's host key against the +// per-profile known-hosts store. An unknown host returns errHostKeyUnknown so +// Java can show the fingerprint and, once confirmed, retry with the key +// trusted; a changed key is rejected outright, as OpenSSH does. When the user +// has confirmed a fingerprint, the callback accepts exactly that key and +// appends it to the store. +func (s *SSHClient) tofuHostKeyCallback() (gossh.HostKeyCallback, error) { + s.mu.Lock() + configDir := s.knownHostsConfigDir + profileID := s.knownHostsProfile + trusted := s.trustHostKey + s.mu.Unlock() + + if configDir == "" || profileID == "" { + return nil, errors.New("no known-hosts store configured for regular SSH") + } + + store, err := openKnownHostsStore(configDir, profileID) + if err != nil { + return nil, fmt.Errorf("load known-hosts store: %w", err) + } + + return func(hostname string, remote net.Addr, key gossh.PublicKey) error { + verdict, err := store.verify(hostname, remote, key) + if err != nil { + return err + } + if verdict == hostKeyMatched { + return nil + } + if verdict == hostKeyChanged { + return fmt.Errorf("SSH host key changed for %s (possible attack)", hostname) + } + + fingerprint := gossh.FingerprintSHA256(key) + if trusted == "" { + return &errHostKeyUnknown{fingerprint: fingerprint} + } + if trusted != fingerprint { + return fmt.Errorf("SSH host key changed since it was confirmed for %s", hostname) + } + if err := store.append(hostname, remote, key); err != nil { + return fmt.Errorf("persist trusted host key: %w", err) + } + // The confirmation is spent: now that the key is stored, a later + // reconnect must verify against the file, not re-accept this fingerprint. + s.mu.Lock() + s.trustHostKey = "" + s.mu.Unlock() + return nil + }, nil +} + +func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config, cfgPath string) (string, error) { + s.mu.Lock() + urlOpener := s.urlOpener + s.mu.Unlock() + if urlOpener == nil { + return "", errors.New("URL opener not configured for JWT auth") + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) + defer cancel() + + flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profileLoginHint(cfgPath)) + if err != nil { + return "", fmt.Errorf("create oauth flow: %w", err) + } + + // The status callback covers the browser round-trip, which would + // otherwise leave the terminal blank. + tokenInfo, err := runOAuthFlow(ctx, flow, urlOpener, func() { + s.notifyStatus("Waiting for browser authentication...") + }) + if err != nil { + return "", err + } + + token := tokenInfo.GetTokenToUse() + if token == "" { + return "", errors.New("empty token returned by IdP") + } + + // Tells the client the browser round-trip is over so it can dismiss the + // surface it opened, the same way the login and session-extend flows do. + // Without it the Custom Tab stays in front of the terminal even though the + // token has already been collected. + urlOpener.OnLoginSuccess() + + return token, nil +} + +func (s *SSHClient) dialAndHandshake(gen uint64, host string, port int, clientConfig *gossh.ClientConfig) error { + addr := net.JoinHostPort(host, strconv.Itoa(port)) + ctx, cancel := context.WithTimeout(context.Background(), sshDialTimeout) + defer cancel() + + s.mu.Lock() + if gen != s.gen { + s.mu.Unlock() + return errClientClosed + } + s.dialCancel = cancel + s.mu.Unlock() + + var dialer net.Dialer + conn, err := dialer.DialContext(ctx, "tcp", addr) + if err != nil { + return fmt.Errorf("dial %s: %w", addr, err) + } + + client, err := nbssh.Handshake(ctx, conn, addr, clientConfig) + if err != nil { + return err + } + + s.mu.Lock() + if gen != s.gen { + s.mu.Unlock() + closeQuiet(client, "stale ssh client") + return errClientClosed + } + s.sshClient = client + listener := s.listener + s.mu.Unlock() + + if listener != nil { + listener.OnConnected() + } + return nil +} + +func (s *SSHClient) readLoop(r io.Reader, name string) string { + buf := make([]byte, 4096) + for { + n, err := r.Read(buf) + if n > 0 { + s.mu.Lock() + listener := s.listener + s.mu.Unlock() + if listener != nil { + chunk := make([]byte, n) + copy(chunk, buf[:n]) + listener.OnData(chunk) + } + } + if err != nil { + // EOF is a normal shell exit, so report it without a reason. + if errors.Is(err, io.EOF) { + return "" + } + log.Debugf("ssh %s read: %v", name, err) + return rootCause(err).Error() + } + } +} + +// notifyStatus writes a progress line to the terminal through the normal +// output path, so long steps are visible while nothing else is arriving. +func (s *SSHClient) notifyStatus(text string) { + s.mu.Lock() + listener := s.listener + s.mu.Unlock() + if listener != nil { + listener.OnData([]byte("\r\n\x1b[33m" + text + "\x1b[0m\r\n")) + } +} + +func (s *SSHClient) notifyClose(gen uint64, reason string) { + s.mu.Lock() + if gen != s.gen || s.closed { + s.mu.Unlock() + return + } + s.closed = true + listener := s.listener + s.mu.Unlock() + if listener != nil { + listener.OnClose(reason) + } +} + +func closeQuiet(c io.Closer, label string) { + if c == nil { + return + } + if err := c.Close(); err != nil && !errors.Is(err, io.EOF) { + log.Debugf("ssh: close %s: %v", label, err) + } +} + +func detectServerType(host string, port int) detection.ServerType { + ctx, cancel := context.WithTimeout(context.Background(), sshDetectionTimeout) + defer cancel() + + dialer := &net.Dialer{} + serverType, err := detection.DetectSSHServerType(ctx, dialer, host, port) + if err != nil { + log.Debugf("ssh: server detection failed: %v (assuming regular SSH)", err) + return detection.ServerTypeRegular + } + return serverType +} + +// rootCause returns the innermost error of a %w chain, so the terminal shows +// "i/o timeout" rather than every layer that added context on the way up. +func rootCause(err error) error { + for { + // A joined error has no single root, so keep it as-is. + if _, ok := err.(interface{ Unwrap() []error }); ok { + return err + } + next := errors.Unwrap(err) + if next == nil { + return err + } + err = next + } +} + +// isAuthFailure distinguishes credential rejection from dial, timeout and +// host-key errors, which retrying with a password would not fix. +func isAuthFailure(err error) bool { + if errors.Is(err, errPasswordRequired) { + return true + } + var partial *gossh.PartialSuccessError + if errors.As(err, &partial) { + return true + } + return strings.Contains(err.Error(), "unable to authenticate") +} + +// passwordCouldHelp reports whether prompting for a password again can change +// the outcome. gossh lists a method under "attempted methods" only when the +// server offered it, so a supplied password that was never attempted means the +// server does not accept passwords and the real error should surface instead. +func passwordCouldHelp(err error, passwordOffered bool) bool { + if !passwordOffered { + return true + } + msg := err.Error() + return strings.Contains(msg, "password") || strings.Contains(msg, "keyboard-interactive") +} diff --git a/client/android/ssh_known_hosts.go b/client/android/ssh_known_hosts.go new file mode 100644 index 000000000..eea90fd32 --- /dev/null +++ b/client/android/ssh_known_hosts.go @@ -0,0 +1,168 @@ +//go:build android + +package android + +import ( + "bytes" + "net" + "strconv" + "strings" + "sync" + + gossh "golang.org/x/crypto/ssh" + "golang.org/x/crypto/ssh/knownhosts" +) + +const knownHostsNamespace = "ssh" + +const ( + hostKeyUnknown hostKeyVerdict = iota + hostKeyMatched + hostKeyChanged +) + +var knownHostsMu sync.Mutex + +type hostKeyVerdict uint8 + +type knownHostsSection struct { + KnownHosts []string `json:"knownHosts"` +} + +type knownHostsStore struct { + prefs prefsStore +} + +// RemoveKnownHost deletes every known-hosts entry for host:port from the +// profile's store, so a host trusted for a session that is being deleted does +// not linger. Java calls this only once no session targets that host, so a +// shared host stays trusted. A missing entry is not an error: the goal state +// is "absent". +func RemoveKnownHost(configDir, profileID, host string, port int) error { + store, err := openKnownHostsStore(configDir, profileID) + if err != nil { + return err + } + return store.removeHost(host, port) +} + +func openKnownHostsStore(configDir, profileID string) (*knownHostsStore, error) { + prefs, err := newProfilePrefs(configDir, profileID) + if err != nil { + return nil, err + } + return &knownHostsStore{prefs: prefs}, nil +} + +func (st *knownHostsStore) verify(hostname string, remote net.Addr, key gossh.PublicKey) (hostKeyVerdict, error) { + lines, err := st.lines() + if err != nil { + return hostKeyUnknown, err + } + targets := knownHostsTargets(hostname, remote) + + verdict := hostKeyUnknown + for _, line := range lines { + pubKey, ok := knownHostsLineKey(line, targets) + if !ok { + continue + } + if pubKey.Type() == key.Type() && bytes.Equal(pubKey.Marshal(), key.Marshal()) { + return hostKeyMatched, nil + } + verdict = hostKeyChanged + } + return verdict, nil +} + +func (st *knownHostsStore) append(hostname string, remote net.Addr, key gossh.PublicKey) error { + line := knownhosts.Line(knownHostsTargets(hostname, remote), key) + + knownHostsMu.Lock() + defer knownHostsMu.Unlock() + + lines, err := st.lines() + if err != nil { + return err + } + return st.prefs.Put(knownHostsNamespace, knownHostsSection{KnownHosts: append(lines, line)}) +} + +func (st *knownHostsStore) removeHost(host string, port int) error { + target := knownhosts.Normalize(net.JoinHostPort(host, strconv.Itoa(port))) + + knownHostsMu.Lock() + defer knownHostsMu.Unlock() + + lines, err := st.lines() + if err != nil { + return err + } + kept := make([]string, 0, len(lines)) + for _, line := range lines { + if knownHostsLineMatches(line, target) { + continue + } + kept = append(kept, line) + } + if len(kept) == len(lines) { + return nil + } + return st.prefs.Put(knownHostsNamespace, knownHostsSection{KnownHosts: kept}) +} + +func (st *knownHostsStore) lines() ([]string, error) { + var section knownHostsSection + if _, err := st.prefs.Get(knownHostsNamespace, §ion); err != nil { + return nil, err + } + return section.KnownHosts, nil +} + +func knownHostsTargets(hostname string, remote net.Addr) []string { + targets := []string{knownhosts.Normalize(hostname)} + if remote != nil { + if normalized := knownhosts.Normalize(remote.String()); normalized != targets[0] { + targets = append(targets, normalized) + } + } + return targets +} + +func knownHostsLineKey(line string, targets []string) (gossh.PublicKey, bool) { + trimmed := strings.TrimSpace(line) + if trimmed == "" || strings.HasPrefix(trimmed, "#") { + return nil, false + } + _, hosts, pubKey, _, _, err := gossh.ParseKnownHosts([]byte(trimmed)) + if err != nil { + return nil, false + } + for _, host := range hosts { + for _, target := range targets { + if host == target { + return pubKey, true + } + } + } + return nil, false +} + +// knownHostsLineMatches reports whether a known-hosts line's address list +// contains the normalized target. Comment and blank lines never match. +func knownHostsLineMatches(line, target string) bool { + trimmed := strings.TrimSpace(line) + if trimmed == "" || strings.HasPrefix(trimmed, "#") { + return false + } + fields := strings.Fields(trimmed) + if len(fields) == 0 { + return false + } + for _, addr := range strings.Split(fields[0], ",") { + if addr == target { + return true + } + } + return false +} diff --git a/client/android/ssh_sessions.go b/client/android/ssh_sessions.go new file mode 100644 index 000000000..44b5464e9 --- /dev/null +++ b/client/android/ssh_sessions.go @@ -0,0 +1,104 @@ +//go:build android + +package android + +const ( + sshSessionsNamespace = "ssh-sessions" + maxStoredSSHSessions = 50 +) + +type sshSessionRecord struct { + ID string `json:"id"` + Host string `json:"host"` + Port int `json:"port"` + User string `json:"user"` +} + +type sshSessionsSection struct { + Sessions []sshSessionRecord `json:"sessions"` +} + +// SSHSessionEntry is one stored SSH session, without any credential. +type SSHSessionEntry struct { + ID string + Host string + Port int + User string +} + +// SSHSessionArray wraps stored SSH sessions for gomobile compatibility. +type SSHSessionArray struct { + items []*SSHSessionEntry +} + +// NewSSHSessionArray creates an empty session array to fill via Add. +func NewSSHSessionArray() *SSHSessionArray { + return &SSHSessionArray{} +} + +// Add appends a session entry, oldest first. +func (a *SSHSessionArray) Add(id, host string, port int, user string) { + a.items = append(a.items, &SSHSessionEntry{ID: id, Host: host, Port: port, User: user}) +} + +// Length returns the number of entries. +func (a *SSHSessionArray) Length() int { + return len(a.items) +} + +// Get returns the entry at index i, or nil when out of range. +func (a *SSHSessionArray) Get(i int) *SSHSessionEntry { + if i < 0 || i >= len(a.items) { + return nil + } + return a.items[i] +} + +// SSHSessionStore reads and writes a profile's stored SSH sessions. +type SSHSessionStore struct { + prefs prefsStore +} + +// NewSSHSessionStore opens the session store of the given profile. +func NewSSHSessionStore(configDir, profileID string) (*SSHSessionStore, error) { + prefs, err := newProfilePrefs(configDir, profileID) + if err != nil { + return nil, err + } + return &SSHSessionStore{prefs: prefs}, nil +} + +// Load returns the stored sessions, oldest first. +func (s *SSHSessionStore) Load() (*SSHSessionArray, error) { + var section sshSessionsSection + if _, err := s.prefs.Get(sshSessionsNamespace, §ion); err != nil { + return nil, err + } + + out := NewSSHSessionArray() + for _, record := range section.Sessions { + if record.ID == "" || record.Host == "" { + continue + } + out.Add(record.ID, record.Host, record.Port, record.User) + } + return out, nil +} + +// Save replaces the stored sessions, keeping only the newest entries when the +// list exceeds the storage cap. +func (s *SSHSessionStore) Save(sessions *SSHSessionArray) error { + var items []*SSHSessionEntry + if sessions != nil { + items = sessions.items + } + if len(items) > maxStoredSSHSessions { + items = items[len(items)-maxStoredSSHSessions:] + } + + records := make([]sshSessionRecord, 0, len(items)) + for _, item := range items { + records = append(records, sshSessionRecord{ID: item.ID, Host: item.Host, Port: item.Port, User: item.User}) + } + return s.prefs.Put(sshSessionsNamespace, sshSessionsSection{Sessions: records}) +} diff --git a/client/anonymize/anonymize.go b/client/anonymize/anonymize.go index c140cef89..c5d43ed55 100644 --- a/client/anonymize/anonymize.go +++ b/client/anonymize/anonymize.go @@ -2,6 +2,7 @@ package anonymize import ( "crypto/rand" + "encoding/base64" "fmt" "math/big" "net" @@ -15,13 +16,88 @@ import ( const anonTLD = ".domain" +// Level selects how much the anonymizer redacts. Levels are ordered: a higher +// level redacts strictly more. On the wire (protos, flags) levels travel as +// their string form. +type Level int + +const ( + // LevelDefault anonymizes public IP addresses, IPv6 ULA, domains, and MAC + // addresses. Internal IPv4 ranges (RFC 1918, CGNAT, link-local) are + // preserved so support can reason about the real topology. + LevelDefault Level = iota + // LevelStrict additionally anonymizes internal IP ranges, peer names, and + // WireGuard public keys. + LevelStrict +) + +// LevelDefaultString and LevelStrictString are the wire forms of the levels, +// for boundaries that pass levels as strings (flags, protos, mobile bindings). +const ( + LevelDefaultString = "default" + LevelStrictString = "strict" +) + +// ParseLevel maps s to a Level. Empty means LevelDefault; anything +// unrecognized maps to LevelStrict so an unknown request never yields less +// anonymization than intended. +func ParseLevel(s string) Level { + switch strings.ToLower(s) { + case "", LevelDefaultString: + return LevelDefault + default: + return LevelStrict + } +} + +// String returns the wire form of the level: "default" or "strict". +func (l Level) String() string { + if l >= LevelStrict { + return LevelStrictString + } + return LevelDefaultString +} + +// protectedDomains are NetBird-operated suffixes that stay recognizable in an +// anonymized bundle. At LevelStrict the labels in front of them (the peer +// name) are still replaced, except under netbird.io, which only hosts +// NetBird infrastructure (api, signal, flow), never peer names. +var protectedDomains = []string{"netbird.io", "netbird.selfhosted", "netbird.cloud", "netbird.stage"} + +const infraDomain = "netbird.io" + +var ( + macColonRegex = regexp.MustCompile(`\b[0-9a-fA-F]{2}(?::[0-9a-fA-F]{2}){5}\b`) + macDashRegex = regexp.MustCompile(`\b[0-9a-fA-F]{2}(?:-[0-9a-fA-F]{2}){5}\b`) + wgKeyRegex = regexp.MustCompile(`\b[A-Za-z0-9+/]{43}=`) +) + type Anonymizer struct { ipAnonymizer map[netip.Addr]netip.Addr domainAnonymizer map[string]string - currentAnonIPv4 netip.Addr - currentAnonIPv6 netip.Addr - startAnonIPv4 netip.Addr - startAnonIPv6 netip.Addr + // domainOrder caches the keys of domainAnonymizer sorted longest-first + // for AnonymizeString; it is rebuilt when the map gains entries. + domainOrder []string + labelAnonymizer map[string]string + labelAnonymized map[string]struct{} + labelCounter uint32 + macAnonymizer map[string]string + macCounter uint32 + wgKeyAnonymizer map[string]string + wgKeyAnonymized map[string]struct{} + currentAnonIPv4 netip.Addr + currentAnonIPv6 netip.Addr + startAnonIPv4 netip.Addr + startAnonIPv6 netip.Addr + + // LevelStrict also anonymizes internal ranges (RFC 1918, CGNAT, + // link-local), replacing them from the dedicated internal pools below so + // a reader can still tell an internal address from a public one. + level Level + currentAnonInternalIPv4 netip.Addr + currentAnonInternalIPv6 netip.Addr + startAnonInternalIPv4 netip.Addr + startAnonInternalIPv6 netip.Addr domainKeyRegex *regexp.Regexp } @@ -32,25 +108,50 @@ func DefaultAddresses() (netip.Addr, netip.Addr) { return netip.AddrFrom4([4]byte{198, 51, 100, 0}), netip.MustParseAddr("2001:db8:ffff::") } +// InternalAddresses returns the pool starts used in strict mode for internal +// ranges. Both are reserved ranges that cannot collide with real addressing: +// 198.18.0.0 (RFC 2544 benchmarking), 2001:db8:1:: (RFC 3849 documentation). +func InternalAddresses() (netip.Addr, netip.Addr) { + return netip.AddrFrom4([4]byte{198, 18, 0, 0}), netip.MustParseAddr("2001:db8:1::") +} + func NewAnonymizer(startIPv4, startIPv6 netip.Addr) *Anonymizer { + internalIPv4, internalIPv6 := InternalAddresses() return &Anonymizer{ ipAnonymizer: map[netip.Addr]netip.Addr{}, domainAnonymizer: map[string]string{}, + labelAnonymizer: map[string]string{}, + labelAnonymized: map[string]struct{}{}, + macAnonymizer: map[string]string{}, + wgKeyAnonymizer: map[string]string{}, + wgKeyAnonymized: map[string]struct{}{}, currentAnonIPv4: startIPv4, currentAnonIPv6: startIPv6, startAnonIPv4: startIPv4, startAnonIPv6: startIPv6, + level: LevelDefault, + currentAnonInternalIPv4: internalIPv4, + currentAnonInternalIPv6: internalIPv6, + startAnonInternalIPv4: internalIPv4, + startAnonInternalIPv6: internalIPv6, + domainKeyRegex: regexp.MustCompile(`\bdomain=([^\s,:"]+)`), } } +// SetLevel selects the anonymization level. The zero value of a new +// Anonymizer is LevelDefault. +func (a *Anonymizer) SetLevel(level Level) { + a.level = level +} + func (a *Anonymizer) AnonymizeIP(ip netip.Addr) netip.Addr { + // Normalize 4-in-6 addresses so ::ffff:192.168.1.1 classifies and maps + // like 192.168.1.1. + ip = ip.Unmap() + if ip.IsLoopback() || - ip.IsLinkLocalUnicast() || - ip.IsLinkLocalMulticast() || - ip.IsInterfaceLocalMulticast() || - (ip.Is4() && ip.IsPrivate()) || ip.IsUnspecified() || ip.IsMulticast() || isWellKnown(ip) || @@ -59,18 +160,100 @@ func (a *Anonymizer) AnonymizeIP(ip netip.Addr) netip.Addr { return ip } + if isInternal(ip) && a.level < LevelStrict { + return ip + } + if _, ok := a.ipAnonymizer[ip]; !ok { - if ip.Is4() { - a.ipAnonymizer[ip] = a.currentAnonIPv4 - a.currentAnonIPv4 = a.currentAnonIPv4.Next() - } else { - a.ipAnonymizer[ip] = a.currentAnonIPv6 - a.currentAnonIPv6 = a.currentAnonIPv6.Next() - } + a.ipAnonymizer[ip] = a.nextAnonIP(ip) } return a.ipAnonymizer[ip] } +func (a *Anonymizer) nextAnonIP(ip netip.Addr) netip.Addr { + // At the strict level, internal addresses (including IPv6 ULA, matched + // by IsPrivate) come from the internal pools so they remain recognizable + // as internal without disclosing the real values. + if a.level >= LevelStrict && (isInternal(ip) || ip.IsPrivate()) { + if ip.Is4() { + anon := a.currentAnonInternalIPv4 + a.currentAnonInternalIPv4 = a.currentAnonInternalIPv4.Next() + return anon + } + anon := a.currentAnonInternalIPv6 + a.currentAnonInternalIPv6 = a.currentAnonInternalIPv6.Next() + return anon + } + + if ip.Is4() { + anon := a.currentAnonIPv4 + a.currentAnonIPv4 = a.currentAnonIPv4.Next() + return anon + } + anon := a.currentAnonIPv6 + a.currentAnonIPv6 = a.currentAnonIPv6.Next() + return anon +} + +// AnonymizeMAC replaces a MAC address with a consistent placeholder from the +// locally administered range starting at 02:00:00:00:00:01, at every +// anonymization level. Broadcast, multicast, all-zero, and already assigned +// placeholder addresses are preserved. The colon and dash spellings of the +// same address share one placeholder; the output keeps the input's separator. +func (a *Anonymizer) AnonymizeMAC(mac string) string { + hw, err := net.ParseMAC(mac) + if err != nil || len(hw) != 6 { + return mac + } + + if isWellKnownMAC(hw) || a.isAnonymizedMAC(hw) { + return mac + } + + key := hw.String() + anon, ok := a.macAnonymizer[key] + if !ok { + a.macCounter++ + anon = fmt.Sprintf("02:00:00:%02x:%02x:%02x", byte(a.macCounter>>16), byte(a.macCounter>>8), byte(a.macCounter)) + a.macAnonymizer[key] = anon + } + + if strings.Contains(mac, "-") { + anon = strings.ReplaceAll(anon, ":", "-") + } + return anon +} + +// isAnonymizedMAC reports whether hw is a placeholder this anonymizer already +// handed out, so a second pass over anonymized output leaves it unchanged. +func (a *Anonymizer) isAnonymizedMAC(hw net.HardwareAddr) bool { + if hw[0] != 0x02 || hw[1] != 0 || hw[2] != 0 { + return false + } + value := uint32(hw[3])<<16 | uint32(hw[4])<<8 | uint32(hw[5]) + return value <= a.macCounter +} + +// AnonymizeWGKey replaces a WireGuard public key with a consistent random +// placeholder of the same shape. Keys are only anonymized at LevelStrict; +// placeholders already handed out pass through unchanged. +func (a *Anonymizer) AnonymizeWGKey(key string) string { + if a.level < LevelStrict || !looksLikeWGKey(key) { + return key + } + if _, ok := a.wgKeyAnonymized[key]; ok { + return key + } + + anon, ok := a.wgKeyAnonymizer[key] + if !ok { + anon = generateAnonymousKey() + a.wgKeyAnonymizer[key] = anon + a.wgKeyAnonymized[anon] = struct{}{} + } + return anon +} + func (a *Anonymizer) AnonymizeUDPAddr(addr net.UDPAddr) net.UDPAddr { // Convert IP to netip.Addr ip, ok := netip.AddrFromSlice(addr.IP) @@ -89,12 +272,12 @@ func (a *Anonymizer) AnonymizeUDPAddr(addr net.UDPAddr) net.UDPAddr { // isInAnonymizedRange checks if an IP is within the range of already assigned anonymized IPs func (a *Anonymizer) isInAnonymizedRange(ip netip.Addr) bool { - if ip.Is4() && ip.Compare(a.startAnonIPv4) >= 0 && ip.Compare(a.currentAnonIPv4) <= 0 { - return true - } else if !ip.Is4() && ip.Compare(a.startAnonIPv6) >= 0 && ip.Compare(a.currentAnonIPv6) <= 0 { - return true + if ip.Is4() { + return inPoolRange(ip, a.startAnonIPv4, a.currentAnonIPv4) || + inPoolRange(ip, a.startAnonInternalIPv4, a.currentAnonInternalIPv4) } - return false + return inPoolRange(ip, a.startAnonIPv6, a.currentAnonIPv6) || + inPoolRange(ip, a.startAnonInternalIPv6, a.currentAnonInternalIPv6) } func (a *Anonymizer) AnonymizeIPString(ip string) string { @@ -118,14 +301,23 @@ func (a *Anonymizer) AnonymizeDomain(domain string) string { baseDomain = domain[:len(domain)-1] } - if strings.HasSuffix(baseDomain, "netbird.io") || - strings.HasSuffix(baseDomain, "netbird.selfhosted") || - strings.HasSuffix(baseDomain, "netbird.cloud") || - strings.HasSuffix(baseDomain, "netbird.stage") || - strings.HasSuffix(baseDomain, anonTLD) { + if strings.HasSuffix(baseDomain, anonTLD) { return domain } + // A reverse zone names an address prefix, so it follows the address rules, + // which also keeps its digit labels intact. + if zone, ok := a.anonymizeReverseZone(baseDomain); ok { + return withTrailingDot(zone, hasDot) + } + + if suffix := protectedSuffix(baseDomain); suffix != "" { + if a.level < LevelStrict || baseDomain == suffix || suffix == infraDomain { + return domain + } + return withTrailingDot(a.anonymizePeerName(baseDomain, suffix), hasDot) + } + parts := strings.Split(baseDomain, ".") if len(parts) < 2 { return domain @@ -141,12 +333,53 @@ func (a *Anonymizer) AnonymizeDomain(domain string) string { } result := strings.Replace(baseDomain, baseForLookup, anonymized, 1) - if hasDot { - result += "." + if a.level >= LevelStrict && len(parts) > 2 { + prefix := strings.TrimSuffix(baseDomain, "."+baseForLookup) + result = a.anonymizeLabels(prefix, "host") + "." + anonymized + // The full mapping feeds AnonymizeString so seeded FQDNs are caught + // in log lines as a whole, labels included. + a.domainAnonymizer[baseDomain] = result + } + return withTrailingDot(result, hasDot) +} + +// anonymizePeerName replaces the labels in front of a protected suffix with +// numbered peer placeholders, keeping the suffix, and records the full +// mapping for string replacement in logs. The numbering keeps a peer +// recognizable across the whole bundle without disclosing its name. +func (a *Anonymizer) anonymizePeerName(baseDomain, suffix string) string { + prefix := strings.TrimSuffix(baseDomain, "."+suffix) + result := a.anonymizeLabels(prefix, "peer") + "." + suffix + if result != baseDomain { + a.domainAnonymizer[baseDomain] = result } return result } +// anonymizeLabels replaces each dot-separated label with a consistent +// numbered placeholder ("-"). Wildcard labels and +// placeholders already handed out pass through unchanged. +func (a *Anonymizer) anonymizeLabels(prefix, placeholder string) string { + labels := strings.Split(prefix, ".") + for i, label := range labels { + if label == "*" { + continue + } + if _, ok := a.labelAnonymized[label]; ok { + continue + } + anon, ok := a.labelAnonymizer[label] + if !ok { + a.labelCounter++ + anon = fmt.Sprintf("%s-%d", placeholder, a.labelCounter) + a.labelAnonymizer[label] = anon + a.labelAnonymized[anon] = struct{}{} + } + labels[i] = anon + } + return strings.Join(labels, ".") +} + func (a *Anonymizer) AnonymizeURI(uri string) string { u, err := url.Parse(uri) if err != nil { @@ -178,17 +411,75 @@ func (a *Anonymizer) AnonymizeString(str string) string { ipv4Regex := regexp.MustCompile(`\b(?:[0-9]{1,3}\.){3}[0-9]{1,3}\b`) ipv6Regex := regexp.MustCompile(`\b([0-9a-fA-F:]+:+[0-9a-fA-F]{0,4})(?:%[0-9a-zA-Z]+)?(?:\/[0-9]{1,3})?(?::[0-9]{1,5})?\b`) + // Reverse zones go first and are then held out of the passes below: their + // labels are digits, which the address patterns would otherwise consume. + str, restoreZones := a.replaceReverseZones(str) + str = ipv4Regex.ReplaceAllStringFunc(str, a.AnonymizeIPString) str = ipv6Regex.ReplaceAllStringFunc(str, a.AnonymizeIPString) - for domain, anonDomain := range a.domainAnonymizer { - str = strings.ReplaceAll(str, domain, anonDomain) + for _, domain := range a.sortedDomains() { + str = strings.ReplaceAll(str, domain, a.domainAnonymizer[domain]) } str = a.AnonymizeSchemeURI(str) str = a.AnonymizeDNSLogLine(str) - return str + // MAC handling runs after the IP passes so preserved IPv6 addresses are + // already out of the way; the separator guard skips matches embedded in a + // longer colon- or dash-separated sequence (such as an IPv6 tail). + str = a.anonymizeMACsInString(str, macColonRegex, ':') + str = a.anonymizeMACsInString(str, macDashRegex, '-') + + if a.level >= LevelStrict { + str = wgKeyRegex.ReplaceAllStringFunc(str, a.AnonymizeWGKey) + } + + return restoreZones(str) +} + +// sortedDomains returns the domain mappings longest-first, so a full-FQDN +// mapping (strict level) is applied before the base-domain mapping it +// contains. The order is rebuilt only when domainAnonymizer has grown. +func (a *Anonymizer) sortedDomains() []string { + if len(a.domainOrder) == len(a.domainAnonymizer) { + return a.domainOrder + } + + a.domainOrder = a.domainOrder[:0] + for domain := range a.domainAnonymizer { + a.domainOrder = append(a.domainOrder, domain) + } + slices.SortFunc(a.domainOrder, func(x, y string) int { + if d := len(y) - len(x); d != 0 { + return d + } + return strings.Compare(x, y) + }) + return a.domainOrder +} + +// anonymizeMACsInString replaces MAC addresses matched by re, skipping +// matches that directly adjoin another sep so a six-group run inside a longer +// separated sequence is left alone. +func (a *Anonymizer) anonymizeMACsInString(str string, re *regexp.Regexp, sep byte) string { + matches := re.FindAllStringIndex(str, -1) + if len(matches) == 0 { + return str + } + + var b strings.Builder + last := 0 + for _, m := range matches { + if (m[0] > 0 && str[m[0]-1] == sep) || (m[1] < len(str) && str[m[1]] == sep) { + continue + } + b.WriteString(str[last:m[0]]) + b.WriteString(a.AnonymizeMAC(str[m[0]:m[1]])) + last = m[1] + } + b.WriteString(str[last:]) + return b.String() } // AnonymizeSchemeURI finds and anonymizes URIs with ws, wss, rel, rels, stun, stuns, turn, and turns schemes. @@ -239,10 +530,79 @@ func isWellKnown(addr netip.Addr) bool { "128.0.0.0", "8000::", // 2nd split subnet for default routes } - if slices.Contains(wellKnown, addr.String()) { + return slices.Contains(wellKnown, addr.String()) +} + +// isInternal reports whether ip identifies a host only within the local +// network: IPv4 private (RFC 1918), CGNAT (RFC 6598), and link-local (v4 and +// v6). These are preserved at the default level so support can reason about +// the real topology, and replaced from the internal pools at the strict +// level. IPv6 ULA is deliberately not internal: its random global ID uniquely +// fingerprints the network, so it is anonymized at every level. +func isInternal(ip netip.Addr) bool { + return (ip.Is4() && ip.IsPrivate()) || + ip.IsLinkLocalUnicast() || + isCGNAT(ip) +} + +func inPoolRange(ip, start, current netip.Addr) bool { + return ip.Compare(start) >= 0 && ip.Compare(current) <= 0 +} + +// isWellKnownMAC reports whether hw carries no stable host identity: all-zero +// or a group address (broadcast and multicast). +func isWellKnownMAC(hw net.HardwareAddr) bool { + if hw[0]&1 == 1 { return true } + for _, b := range hw { + if b != 0 { + return false + } + } + return true +} +// looksLikeWGKey reports whether s has the shape of a WireGuard key: +// 44 base64 characters decoding to 32 bytes. +func looksLikeWGKey(s string) bool { + if len(s) != 44 || s[43] != '=' { + return false + } + decoded, err := base64.StdEncoding.DecodeString(s) + return err == nil && len(decoded) == 32 +} + +func generateAnonymousKey() string { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return strings.Repeat("A", 43) + "=" + } + return base64.StdEncoding.EncodeToString(buf) +} + +// protectedSuffix returns the protected NetBird suffix baseDomain ends with, +// or empty. The match is label-anchored so an unrelated domain that merely +// ends in the same characters is not preserved. +func protectedSuffix(baseDomain string) string { + for _, d := range protectedDomains { + if baseDomain == d || strings.HasSuffix(baseDomain, "."+d) { + return d + } + } + return "" +} + +func withTrailingDot(domain string, hasDot bool) string { + if hasDot { + return domain + "." + } + return domain +} + +// isCGNAT reports whether addr is in 100.64.0.0/10 (RFC 6598), the range +// NetBird assigns overlay peer addresses from. +func isCGNAT(addr netip.Addr) bool { cgnatRangeStart := netip.AddrFrom4([4]byte{100, 64, 0, 0}) cgnatRange := netip.PrefixFrom(cgnatRangeStart, 10) diff --git a/client/anonymize/anonymize_test.go b/client/anonymize/anonymize_test.go index 852315fa1..7c3c7bcf8 100644 --- a/client/anonymize/anonymize_test.go +++ b/client/anonymize/anonymize_test.go @@ -1,8 +1,11 @@ package anonymize_test import ( + "bytes" + "encoding/base64" "net/netip" "regexp" + "strings" "testing" "github.com/stretchr/testify/assert" @@ -44,6 +47,301 @@ func TestAnonymizeIP(t *testing.T) { } } +func TestParseLevel(t *testing.T) { + tests := []struct { + input string + expect anonymize.Level + }{ + {"", anonymize.LevelDefault}, + {"default", anonymize.LevelDefault}, + {"DEFAULT", anonymize.LevelDefault}, + {"strict", anonymize.LevelStrict}, + {"STRICT", anonymize.LevelStrict}, + // Unknown values must never yield less anonymization than requested. + {"garbage", anonymize.LevelStrict}, + } + + for _, tc := range tests { + t.Run("input="+tc.input, func(t *testing.T) { + assert.Equal(t, tc.expect, anonymize.ParseLevel(tc.input), "parsed level should match") + }) + } +} + +func TestAnonymizeIP_DefaultLevelInternalRanges(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + + tests := []struct { + name string + ip string + expect string + }{ + {"RFC1918 10/8", "10.1.2.3", "10.1.2.3"}, + {"RFC1918 172.16/12", "172.16.5.5", "172.16.5.5"}, + {"RFC1918 192.168/16", "192.168.1.1", "192.168.1.1"}, + {"CGNAT", "100.64.0.5", "100.64.0.5"}, + {"IPv4 link-local", "169.254.1.1", "169.254.1.1"}, + {"IPv6 link-local", "fe80::1", "fe80::1"}, + // ULA is anonymized even at the default level: its random global ID + // uniquely fingerprints the network, unlike shared RFC 1918 space. + {"IPv6 ULA", "fd12:3456:789a::1", "2001:db8:ffff::"}, + // 4-in-6 addresses classify like their unmapped IPv4 form. + {"4-in-6 RFC1918", "::ffff:192.168.1.1", "192.168.1.1"}, + {"4-in-6 CGNAT", "::ffff:100.64.0.5", "100.64.0.5"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + result := anonymizer.AnonymizeIP(netip.MustParseAddr(tc.ip)) + assert.Equal(t, tc.expect, result.String(), "default level should preserve internal ranges except ULA") + }) + } +} + +func TestAnonymizeIP_StrictLevel(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(anonymize.LevelStrict) + + // Order matters: internal pool addresses are assigned sequentially. + tests := []struct { + name string + ip string + expect string + }{ + {"RFC1918 192.168/16", "192.168.1.1", "198.18.0.0"}, + {"Second RFC1918", "192.168.1.2", "198.18.0.1"}, + {"Repeated RFC1918", "192.168.1.1", "198.18.0.0"}, + {"RFC1918 10/8", "10.1.2.3", "198.18.0.2"}, + {"RFC1918 172.16/12", "172.16.5.5", "198.18.0.3"}, + {"CGNAT", "100.64.0.5", "198.18.0.4"}, + {"IPv4 link-local", "169.254.1.1", "198.18.0.5"}, + {"Public IPv4 uses public pool", "1.2.3.4", "198.51.100.0"}, + {"IPv6 link-local", "fe80::1", "2001:db8:1::"}, + {"IPv6 ULA", "fd12:3456:789a::1", "2001:db8:1::1"}, + {"Public IPv6 uses public pool", "2607:f8b0:4005:805::200e", "2001:db8:ffff::"}, + {"Loopback IPv4", "127.0.0.1", "127.0.0.1"}, + {"Loopback IPv6", "::1", "::1"}, + {"Unspecified", "0.0.0.0", "0.0.0.0"}, + {"Multicast", "224.0.0.251", "224.0.0.251"}, + {"Well known resolver", "8.8.8.8", "8.8.8.8"}, + {"Well known split marker", "128.0.0.0", "128.0.0.0"}, + {"In internal pool range", "198.18.0.3", "198.18.0.3"}, + {"In public pool range", "198.51.100.0", "198.51.100.0"}, + {"4-in-6 repeated RFC1918", "::ffff:192.168.1.1", "198.18.0.0"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + result := anonymizer.AnonymizeIP(netip.MustParseAddr(tc.ip)) + assert.Equal(t, tc.expect, result.String(), "strict level should replace internal ranges from the internal pools") + }) + } +} + +func TestAnonymizeString_StrictInternalIPs(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(anonymize.LevelStrict) + + input := "route 10.20.30.0/24 via 192.168.1.1 dev eth0 src 100.64.0.7" + firstPass := anonymizer.AnonymizeString(input) + secondPass := anonymizer.AnonymizeString(firstPass) + + assert.NotContains(t, firstPass, "10.20.30.0", "private network address should be anonymized") + assert.NotContains(t, firstPass, "192.168.1.1", "private gateway should be anonymized") + assert.NotContains(t, firstPass, "100.64.0.7", "CGNAT address should be anonymized") + assert.Contains(t, firstPass, "/24", "prefix length should be preserved") + assert.Equal(t, firstPass, secondPass, "second pass should not further anonymize the string") +} + +func TestAnonymizeMAC(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + + first := anonymizer.AnonymizeMAC("aa:bb:cc:dd:ee:0f") + assert.Equal(t, "02:00:00:00:00:01", first, "first MAC should get the first placeholder") + assert.Equal(t, first, anonymizer.AnonymizeMAC("aa:bb:cc:dd:ee:0f"), "repeated MAC should map to the same placeholder") + assert.Equal(t, first, anonymizer.AnonymizeMAC("AA:BB:CC:DD:EE:0F"), "case should not affect the mapping") + assert.Equal(t, "02-00-00-00-00-01", anonymizer.AnonymizeMAC("AA-BB-CC-DD-EE-0F"), "dash form should keep its separator but share the mapping") + + second := anonymizer.AnonymizeMAC("10:22:33:44:55:66") + assert.Equal(t, "02:00:00:00:00:02", second, "second distinct MAC should get the next placeholder") + + tests := []struct { + name string + mac string + }{ + {"Broadcast", "ff:ff:ff:ff:ff:ff"}, + {"IPv4 multicast", "01:00:5e:00:00:fb"}, + {"IPv6 multicast", "33:33:00:00:00:01"}, + {"All zero", "00:00:00:00:00:00"}, + {"Assigned placeholder", "02:00:00:00:00:01"}, + {"Invalid", "not-a-mac"}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.mac, anonymizer.AnonymizeMAC(tc.mac), "should be preserved") + }) + } +} + +func TestAnonymizeString_MACAddresses(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + + tests := []struct { + name string + input string + expect string + }{ + { + name: "nftables ether rule", + input: "ether saddr aa:bb:cc:dd:ee:ff drop", + expect: "ether saddr 02:00:00:00:00:01 drop", + }, + { + name: "Windows dash form", + input: "Physical Address : AA-BB-CC-DD-EE-FF", + expect: "Physical Address : 02-00-00-00-00-01", + }, + { + name: "IPv6 address tail is not treated as MAC", + input: "addr fe80:0:11:22:33:44:55:66 scope link", + expect: "addr fe80:0:11:22:33:44:55:66 scope link", + }, + { + name: "broadcast MAC preserved", + input: "dst ff:ff:ff:ff:ff:ff type ARP", + expect: "dst ff:ff:ff:ff:ff:ff type ARP", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + result := anonymizer.AnonymizeString(tc.input) + assert.Equal(t, tc.expect, result, "MAC addresses should be anonymized at every level") + assert.Equal(t, result, anonymizer.AnonymizeString(result), "second pass should not change the result") + }) + } +} + +func TestAnonymizeWGKey(t *testing.T) { + key := base64.StdEncoding.EncodeToString(bytes.Repeat([]byte{0x42}, 32)) + + t.Run("default level preserves keys", func(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + assert.Equal(t, key, anonymizer.AnonymizeWGKey(key), "default level should not touch WireGuard keys") + }) + + t.Run("strict level replaces keys", func(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(anonymize.LevelStrict) + + anon := anonymizer.AnonymizeWGKey(key) + assert.NotEqual(t, key, anon, "strict level should replace the key") + assert.Regexp(t, `^[A-Za-z0-9+/]{43}=$`, anon, "placeholder should keep the WireGuard key shape") + assert.Equal(t, anon, anonymizer.AnonymizeWGKey(key), "repeated key should map to the same placeholder") + assert.Equal(t, anon, anonymizer.AnonymizeWGKey(anon), "an assigned placeholder should pass through unchanged") + + assert.Equal(t, "not-a-key", anonymizer.AnonymizeWGKey("not-a-key"), "non-key values should be preserved") + }) +} + +func TestAnonymizeString_WGKeys(t *testing.T) { + key := base64.StdEncoding.EncodeToString(bytes.Repeat([]byte{0x42}, 32)) + input := "peer " + key + " handshake completed" + + t.Run("default level preserves keys", func(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + assert.Equal(t, input, anonymizer.AnonymizeString(input), "default level should not touch WireGuard keys in strings") + }) + + t.Run("strict level replaces keys", func(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(anonymize.LevelStrict) + + firstPass := anonymizer.AnonymizeString(input) + assert.NotContains(t, firstPass, key, "the key should not survive strict anonymization") + assert.Equal(t, anonymizer.AnonymizeWGKey(key), extractKey(t, firstPass), "string replacement should be consistent with AnonymizeWGKey") + assert.Equal(t, firstPass, anonymizer.AnonymizeString(firstPass), "second pass should not change the result") + }) +} + +func extractKey(t *testing.T, logLine string) string { + t.Helper() + fields := strings.Fields(logLine) + require.Len(t, fields, 4, "log line should keep its structure") + return fields[1] +} + +func TestAnonymizeDomain_StrictLevel(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(anonymize.LevelStrict) + + t.Run("netbird peer name", func(t *testing.T) { + result := anonymizer.AnonymizeDomain("my-laptop.netbird.cloud") + assert.Regexp(t, `^peer-\d+\.netbird\.cloud$`, result, "peer name should be anonymized, suffix kept") + assert.NotContains(t, result, "my-laptop", "the peer name should not survive") + assert.Equal(t, result, anonymizer.AnonymizeDomain("my-laptop.netbird.cloud"), "repeated domain should map consistently") + assert.Equal(t, result, anonymizer.AnonymizeDomain(result), "an anonymized domain should pass through unchanged") + }) + + t.Run("bare netbird domain", func(t *testing.T) { + assert.Equal(t, "netbird.cloud", anonymizer.AnonymizeDomain("netbird.cloud"), "the bare protected suffix should be preserved") + }) + + t.Run("netbird infrastructure preserved", func(t *testing.T) { + assert.Equal(t, "api.netbird.io", anonymizer.AnonymizeDomain("api.netbird.io"), + "netbird.io hosts infrastructure, not peer names, and should stay readable") + }) + + t.Run("leading labels of other domains", func(t *testing.T) { + result := anonymizer.AnonymizeDomain("host1.corp.example.com") + assert.Regexp(t, `^host-\d+\.host-\d+\.anon-[a-zA-Z0-9]+\.domain$`, result, "every label should be anonymized") + for _, label := range []string{"host1", "corp", "example"} { + assert.NotContains(t, result, label, "no original label should survive") + } + assert.Equal(t, result, anonymizer.AnonymizeDomain("host1.corp.example.com"), "repeated domain should map consistently") + }) + + t.Run("same label maps consistently across domains", func(t *testing.T) { + first := anonymizer.AnonymizeDomain("shared.one.com") + second := anonymizer.AnonymizeDomain("shared.two.com") + assert.Equal(t, strings.Split(first, ".")[0], strings.Split(second, ".")[0], "the shared host label should get one placeholder") + }) + + t.Run("wildcard label preserved", func(t *testing.T) { + result := anonymizer.AnonymizeDomain("*.example.com") + assert.Regexp(t, `^\*\.anon-[a-zA-Z0-9]+\.domain$`, result, "the wildcard label should stay a wildcard") + }) +} + +func TestAnonymizeDomain_DefaultLevelKeepsPeerNames(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + + assert.Equal(t, "my-laptop.netbird.cloud", anonymizer.AnonymizeDomain("my-laptop.netbird.cloud"), + "default level should preserve netbird FQDNs including the peer name") + assert.Regexp(t, `^sub\.anon-[a-zA-Z0-9]+\.domain$`, anonymizer.AnonymizeDomain("sub.example.com"), + "default level should keep subdomain labels") +} + +func TestAnonymizeString_StrictPeerNames(t *testing.T) { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(anonymize.LevelStrict) + + // Seed like the bundle generator does from the status: base first, then + // the full FQDN, so replacement must prefer the longer mapping. + anonBase := anonymizer.AnonymizeDomain("example.com") + anonPeer := anonymizer.AnonymizeDomain("peer1.netbird.cloud") + anonHost := anonymizer.AnonymizeDomain("host1.example.com") + + logLine := "connected to peer1.netbird.cloud via host1.example.com endpoint" + firstPass := anonymizer.AnonymizeString(logLine) + assert.NotContains(t, firstPass, "peer1", "the peer name should not survive in logs") + assert.NotContains(t, firstPass, "host1", "the host label should not survive in logs") + assert.Contains(t, firstPass, anonPeer, "the seeded peer mapping should be applied") + assert.Contains(t, firstPass, anonHost, "the seeded host mapping should be applied, not just the base mapping") + assert.NotContains(t, firstPass, "host1."+anonBase, "the base mapping must not preempt the longer FQDN mapping") + assert.Equal(t, firstPass, anonymizer.AnonymizeString(firstPass), "second pass should not change the result") +} + func TestAnonymizeDNSLogLine(t *testing.T) { anonymizer := anonymize.NewAnonymizer(netip.Addr{}, netip.Addr{}) tests := []struct { diff --git a/client/anonymize/reverse_zone.go b/client/anonymize/reverse_zone.go new file mode 100644 index 000000000..b521b71b7 --- /dev/null +++ b/client/anonymize/reverse_zone.go @@ -0,0 +1,174 @@ +package anonymize + +import ( + "encoding/hex" + "net/netip" + "regexp" + "strconv" + "strings" +) + +const ( + reverseZoneSuffixV4 = ".in-addr.arpa" + reverseZoneSuffixV6 = ".ip6.arpa" + + v6Nibbles = 32 + v4Octets = 4 +) + +// reverseZoneRegexes match a reverse zone or a full reverse name in free text. +// They are applied before the address passes of AnonymizeString, whose IPv4 +// pattern would otherwise consume the digit labels of a zone and replace parts +// of it with unrelated addresses. +var reverseZoneRegexes = []*regexp.Regexp{ + regexp.MustCompile(`(?:[0-9]{1,3}\.){1,4}in-addr\.arpa\b`), + regexp.MustCompile(`(?:[0-9a-fA-F]\.){1,32}ip6\.arpa\b`), +} + +// anonymizeReverseZone maps a reverse zone to the zone of the anonymized form +// of the prefix it encodes, so it follows the address rules rather than the +// domain ones: the zone of an address that is preserved is preserved too, and +// the zone of one that is replaced names the replacement. This keeps a reverse +// zone recognizable as such, and consistent with the addresses it belongs to +// elsewhere in the same output. It reports false for anything that is not a +// reverse zone. +func (a *Anonymizer) anonymizeReverseZone(domain string) (string, bool) { + prefix, labelCount, suffix, ok := parseReverseZone(domain) + if !ok { + return "", false + } + + anonymized := a.AnonymizeIP(prefix) + if anonymized == prefix { + return domain, true + } + + return reverseZoneName(anonymized, labelCount) + suffix, true +} + +// replaceReverseZones anonymizes every reverse zone in str and swaps each one +// for a placeholder, returning a function that puts the anonymized zones back. +// The placeholders carry no dots, digits or colons, so no later pass matches +// them. +func (a *Anonymizer) replaceReverseZones(str string) (string, func(string) string) { + var zones []string + + for _, re := range reverseZoneRegexes { + str = re.ReplaceAllStringFunc(str, func(match string) string { + zone, ok := a.anonymizeReverseZone(match) + if !ok { + return match + } + + zones = append(zones, zone) + return reverseZonePlaceholder(len(zones) - 1) + }) + } + + if len(zones) == 0 { + return str, func(s string) string { return s } + } + + return str, func(s string) string { + for i, zone := range zones { + s = strings.ReplaceAll(s, reverseZonePlaceholder(i), zone) + } + return s + } +} + +func reverseZonePlaceholder(index int) string { + return "\x00reversezone" + strconv.Itoa(index) + "\x00" +} + +// parseReverseZone turns a reverse zone into the address of the prefix its +// labels spell backwards, padding the absent low-order part with zeroes, and +// returns the label count and zone suffix so the name can be rebuilt. +func parseReverseZone(domain string) (netip.Addr, int, string, bool) { + lower := strings.ToLower(domain) + + switch { + case strings.HasSuffix(lower, reverseZoneSuffixV4): + labels := strings.Split(strings.TrimSuffix(lower, reverseZoneSuffixV4), ".") + addr, ok := reverseZoneAddrV4(labels) + return addr, len(labels), reverseZoneSuffixV4, ok + case strings.HasSuffix(lower, reverseZoneSuffixV6): + labels := strings.Split(strings.TrimSuffix(lower, reverseZoneSuffixV6), ".") + addr, ok := reverseZoneAddrV6(labels) + return addr, len(labels), reverseZoneSuffixV6, ok + default: + return netip.Addr{}, 0, "", false + } +} + +func reverseZoneAddrV4(labels []string) (netip.Addr, bool) { + if len(labels) == 0 || len(labels) > v4Octets { + return netip.Addr{}, false + } + + var octets [v4Octets]byte + for i, label := range labels { + octet, err := strconv.ParseUint(label, 10, 8) + if err != nil { + return netip.Addr{}, false + } + octets[len(labels)-1-i] = byte(octet) + } + + return netip.AddrFrom4(octets), true +} + +func reverseZoneAddrV6(labels []string) (netip.Addr, bool) { + if len(labels) == 0 || len(labels) > v6Nibbles { + return netip.Addr{}, false + } + + nibbles := make([]byte, 0, v6Nibbles) + for i := len(labels) - 1; i >= 0; i-- { + if len(labels[i]) != 1 || !isHexDigit(labels[i][0]) { + return netip.Addr{}, false + } + nibbles = append(nibbles, labels[i][0]) + } + for len(nibbles) < v6Nibbles { + nibbles = append(nibbles, '0') + } + + var groups []string + for i := 0; i < len(nibbles); i += 4 { + groups = append(groups, string(nibbles[i:i+4])) + } + + addr, err := netip.ParseAddr(strings.Join(groups, ":")) + if err != nil { + return netip.Addr{}, false + } + + return addr, true +} + +// reverseZoneName spells the first labelCount labels of addr backwards, the +// inverse of parseReverseZone, without the zone suffix. +func reverseZoneName(addr netip.Addr, labelCount int) string { + labels := make([]string, 0, labelCount) + + if addr.Is4() { + octets := addr.As4() + for i := labelCount - 1; i >= 0; i-- { + labels = append(labels, strconv.Itoa(int(octets[i]))) + } + return strings.Join(labels, ".") + } + + address := addr.As16() + nibbles := hex.EncodeToString(address[:]) + for i := labelCount - 1; i >= 0; i-- { + labels = append(labels, string(nibbles[i])) + } + + return strings.Join(labels, ".") +} + +func isHexDigit(c byte) bool { + return c >= '0' && c <= '9' || c >= 'a' && c <= 'f' || c >= 'A' && c <= 'F' +} diff --git a/client/anonymize/reverse_zone_test.go b/client/anonymize/reverse_zone_test.go new file mode 100644 index 000000000..8c3b8954a --- /dev/null +++ b/client/anonymize/reverse_zone_test.go @@ -0,0 +1,171 @@ +package anonymize + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func newLeveledAnonymizer(level Level) *Anonymizer { + a := NewAnonymizer(DefaultAddresses()) + a.SetLevel(level) + return a +} + +// TestAnonymizeDomainReverseZone covers reverse zones going through the address +// rules instead of the domain ones, so a zone stays a zone and an address that +// is preserved keeps the zone that names it. +func TestAnonymizeDomainReverseZone(t *testing.T) { + // 100.64.0.0/10 is the overlay range, which is CGNAT: preserved at the + // default level and replaced from the internal pool at the strict one + const overlayZone = "64.100.in-addr.arpa" + + t.Run("overlay zone preserved at the default level", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + assert.Equal(t, overlayZone, a.AnonymizeDomain(overlayZone), "should keep the zone of a preserved address") + }) + + t.Run("private zone preserved at the default level", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + assert.Equal(t, "168.192.in-addr.arpa", a.AnonymizeDomain("168.192.in-addr.arpa"), "should keep the zone of a private address") + }) + + t.Run("overlay zone replaced at the strict level", func(t *testing.T) { + a := newLeveledAnonymizer(LevelStrict) + + got := a.AnonymizeDomain(overlayZone) + require.True(t, strings.HasSuffix(got, reverseZoneSuffixV4), "should stay a reverse zone, got %q", got) + assert.NotEqual(t, overlayZone, got, "should replace the encoded prefix") + assert.Len(t, strings.Split(strings.TrimSuffix(got, reverseZoneSuffixV4), "."), 2, + "should keep the label count, got %q", got) + }) + + t.Run("public zone replaced at the default level", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + got := a.AnonymizeDomain("113.0.203.in-addr.arpa") + require.True(t, strings.HasSuffix(got, reverseZoneSuffixV4), "should stay a reverse zone, got %q", got) + assert.NotEqual(t, "113.0.203.in-addr.arpa", got, "should replace a public prefix") + }) + + t.Run("zone of an address keeps that address mapping", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + anonymizedAddr := a.AnonymizeIPString("203.0.113.7") + got := a.AnonymizeDomain("7.113.0.203.in-addr.arpa") + + octets := strings.Split(anonymizedAddr, ".") + want := octets[3] + "." + octets[2] + "." + octets[1] + "." + octets[0] + reverseZoneSuffixV4 + assert.Equal(t, want, got, "should name the same replacement as the address itself") + }) + + t.Run("ipv6 nibble labels stay single digits", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + zone := "0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.2.0.0.0" + reverseZoneSuffixV6 + got := a.AnonymizeDomain(zone) + + require.True(t, strings.HasSuffix(got, reverseZoneSuffixV6), "should stay a reverse zone, got %q", got) + labels := strings.Split(strings.TrimSuffix(got, reverseZoneSuffixV6), ".") + assert.Len(t, labels, 28, "should keep every nibble label, got %q", got) + for _, label := range labels { + assert.Len(t, label, 1, "nibble label %q should stay a single digit", label) + } + }) + + t.Run("trailing dot is kept", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + assert.Equal(t, "64.100.in-addr.arpa.", a.AnonymizeDomain("64.100.in-addr.arpa."), "should keep the trailing dot") + }) + + t.Run("a domain that only looks like a zone is anonymized as a domain", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + got := a.AnonymizeDomain("not-a-zone.in-addr.arpa") + assert.NotContains(t, got, "in-addr.arpa", "should fall back to domain anonymization") + }) +} + +// TestAnonymizeStringReverseZone verifies that a zone inside free text, such as +// a DNS log line, is not chewed up by the address passes. The IPv4 pattern +// matches any run of dotted digits, which a reverse zone is made of. +func TestAnonymizeStringReverseZone(t *testing.T) { + t.Run("ipv6 zone survives the address passes", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + zone := "0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.2.0.0.0" + reverseZoneSuffixV6 + got := a.AnonymizeString("question: domain=" + zone + " type=PTR") + + assert.Contains(t, got, "type=PTR", "should keep the rest of the line") + assert.NotContains(t, got, "198.51.100", "should not rewrite nibble labels as an address") + + labels := strings.Split(strings.TrimSuffix(strings.TrimPrefix(got, "question: domain="), reverseZoneSuffixV6+" type=PTR"), ".") + assert.Len(t, labels, 28, "should keep every nibble label, got %q", got) + }) + + t.Run("preserved ipv4 zone is untouched", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + line := "reverse zone 64.100.in-addr.arpa registered" + assert.Equal(t, line, a.AnonymizeString(line), "should keep the zone of a preserved address") + }) + + t.Run("public ipv4 zone is replaced consistently", func(t *testing.T) { + a := newLeveledAnonymizer(LevelDefault) + + got := a.AnonymizeString("zone 113.0.203.in-addr.arpa and address 203.0.113.7") + assert.NotContains(t, got, "113.0.203.in-addr.arpa", "should replace the zone") + assert.NotContains(t, got, "203.0.113.7", "should replace the address") + assert.Contains(t, got, reverseZoneSuffixV4, "should keep the zone suffix") + }) +} + +func TestParseReverseZone(t *testing.T) { + tests := []struct { + name string + zone string + addr string + labels int + }{ + {name: "v4 two labels", zone: "0.100" + reverseZoneSuffixV4, addr: "100.0.0.0", labels: 2}, + {name: "v4 three labels", zone: "1.168.192" + reverseZoneSuffixV4, addr: "192.168.1.0", labels: 3}, + {name: "v4 full address", zone: "7.113.0.203" + reverseZoneSuffixV4, addr: "203.0.113.7", labels: 4}, + { + name: "v6 prefix", + zone: "0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.2.0.0.0" + reverseZoneSuffixV6, + addr: "2::", + labels: 28, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + addr, labels, suffix, ok := parseReverseZone(tc.zone) + require.True(t, ok, "should decode the reverse zone") + assert.Equal(t, tc.addr, addr.String(), "should decode to the encoded prefix") + assert.Equal(t, tc.labels, labels, "should count the labels") + assert.Equal(t, tc.zone, reverseZoneName(addr, labels)+suffix, "should re-encode to the original zone") + }) + } +} + +func TestParseReverseZoneRejectsNonZones(t *testing.T) { + tests := []string{ + "example.com", + "in-addr.arpa", + "x.100" + reverseZoneSuffixV4, + "256" + reverseZoneSuffixV4, + "1.2.3.4.5" + reverseZoneSuffixV4, + "ab" + reverseZoneSuffixV6, + "g" + reverseZoneSuffixV6, + } + + for _, zone := range tests { + t.Run(zone, func(t *testing.T) { + _, _, _, ok := parseReverseZone(zone) + assert.False(t, ok, "should reject %q", zone) + }) + } +} diff --git a/client/cmd/daemon_error.go b/client/cmd/daemon_error.go new file mode 100644 index 000000000..0d5b1307e --- /dev/null +++ b/client/cmd/daemon_error.go @@ -0,0 +1,66 @@ +package cmd + +import ( + "errors" + "fmt" + "strings" + + "google.golang.org/genproto/googleapis/rpc/errdetails" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// daemonCallError prepares a daemon error for display. A refusal the daemon +// raised because the operation needs root/administrator is already guidance +// written for the user, so it is surfaced on its own instead of buried under the +// gRPC envelope and the name of the RPC that hit it. Anything else is wrapped +// with context as usual. +func daemonCallError(context string, err error) error { + if guidance, ok := privilegeGuidance(err); ok { + return errors.New(guidance) + } + return fmt.Errorf("%s: %w", context, err) +} + +// privilegeGuidance renders the daemon's privilege refusal as a summary and the +// command that performs the operation with the privileges it needs. It reports +// false for any other error. +func privilegeGuidance(err error) (string, bool) { + info, ok := privilegeErrorInfo(err) + if !ok { + return "", false + } + + summary := info.GetMetadata()[ipcauth.ErrorMetaSummary] + command := info.GetMetadata()[ipcauth.ErrorMetaCommand] + if summary == "" { + // Detail without a summary: fall back to the status message, which + // carries the same text. + summary = strings.TrimSpace(gstatus.Convert(err).Message()) + } + if command == "" { + return summary, true + } + + return fmt.Sprintf("%s\n\n %s\n", summary, command), true +} + +// privilegeErrorInfo returns the daemon's privilege-refusal detail, if the error +// carries one. +func privilegeErrorInfo(err error) (*errdetails.ErrorInfo, bool) { + if err == nil { + return nil, false + } + + for _, detail := range gstatus.Convert(err).Details() { + info, ok := detail.(*errdetails.ErrorInfo) + if !ok { + continue + } + if info.GetReason() == ipcauth.ErrorReasonPrivilegeRequired && info.GetDomain() == ipcauth.ErrorDomain { + return info, true + } + } + return nil, false +} diff --git a/client/cmd/debug.go b/client/cmd/debug.go index 57e75f663..893b1e248 100644 --- a/client/cmd/debug.go +++ b/client/cmd/debug.go @@ -27,10 +27,11 @@ import ( const errCloseConnection = "Failed to close connection: %v" var ( - logFileCount uint32 - systemInfoFlag bool - uploadBundleFlag bool - uploadBundleURLFlag string + logFileCount uint32 + systemInfoFlag bool + uploadBundleFlag bool + uploadBundleURLFlag string + uploadBundleInsecureFlag bool ) var debugCmd = &cobra.Command{ @@ -155,6 +156,11 @@ func debugConfigDump(cmd *cobra.Command, _ []string) error { // request. Returns an error if the RPC fails or if the daemon reports // an upload failure reason. func debugBundle(cmd *cobra.Command, _ []string) error { + anonymizeEnabled, anonymizeLevel, err := effectiveAnonymize() + if err != nil { + return err + } + conn, err := getClient(cmd) if err != nil { return err @@ -167,17 +173,19 @@ func debugBundle(cmd *cobra.Command, _ []string) error { client := proto.NewDaemonServiceClient(conn) request := &proto.DebugBundleRequest{ - Anonymize: anonymizeFlag, - SystemInfo: systemInfoFlag, - LogFileCount: logFileCount, - CliVersion: version.NetbirdVersion(), + Anonymize: anonymizeEnabled, + AnonymizeLevel: anonymizeLevel.String(), + SystemInfo: systemInfoFlag, + LogFileCount: logFileCount, + CliVersion: version.NetbirdVersion(), } if uploadBundleFlag { request.UploadURL = uploadBundleURLFlag + request.UploadInsecure = uploadBundleInsecureFlag } resp, err := client.DebugBundle(cmd.Context(), request) if err != nil { - return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message()) + return daemonCallError("bundle debug", err) } cmd.Printf("Local file:\n%s\n", resp.GetPath()) @@ -227,6 +235,11 @@ func runForDuration(cmd *cobra.Command, args []string) error { return fmt.Errorf("invalid duration format: %v", err) } + anonymizeEnabled, anonymizeLevel, err := effectiveAnonymize() + if err != nil { + return err + } + conn, err := getClient(cmd) if err != nil { return err @@ -366,17 +379,19 @@ func runForDuration(cmd *cobra.Command, args []string) error { cmd.Println("Creating debug bundle...") request := &proto.DebugBundleRequest{ - Anonymize: anonymizeFlag, - SystemInfo: systemInfoFlag, - LogFileCount: logFileCount, - CliVersion: version.NetbirdVersion(), + Anonymize: anonymizeEnabled, + AnonymizeLevel: anonymizeLevel.String(), + SystemInfo: systemInfoFlag, + LogFileCount: logFileCount, + CliVersion: version.NetbirdVersion(), } if uploadBundleFlag { request.UploadURL = uploadBundleURLFlag + request.UploadInsecure = uploadBundleInsecureFlag } resp, err := client.DebugBundle(cmd.Context(), request) if err != nil { - return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message()) + return daemonCallError("bundle debug", err) } if needsRestoreUp { @@ -524,10 +539,12 @@ func init() { debugBundleCmd.Flags().BoolVarP(&systemInfoFlag, "system-info", "S", true, "Adds system information to the debug bundle") debugBundleCmd.Flags().BoolVarP(&uploadBundleFlag, "upload-bundle", "U", false, "Uploads the debug bundle to a server") debugBundleCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle") + debugBundleCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root") forCmd.Flags().Uint32VarP(&logFileCount, "log-file-count", "C", 1, "Number of rotated log files to include in debug bundle") forCmd.Flags().BoolVarP(&systemInfoFlag, "system-info", "S", true, "Adds system information to the debug bundle") forCmd.Flags().BoolVarP(&uploadBundleFlag, "upload-bundle", "U", false, "Uploads the debug bundle to a server") 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") } diff --git a/client/cmd/login.go b/client/cmd/login.go index a53cb6d5f..6aa019896 100644 --- a/client/cmd/login.go +++ b/client/cmd/login.go @@ -5,7 +5,6 @@ import ( "fmt" "os" "os/user" - "runtime" "strings" log "github.com/sirupsen/logrus" @@ -121,7 +120,7 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str loginRequest := proto.LoginRequest{ SetupKey: providedSetupKey, ManagementUrl: managementURL, - IsUnixDesktopClient: isUnixRunningDesktop(), + IsUnixDesktopClient: util.HasGraphicalSession(), Hostname: hostName, DnsLabels: dnsLabelsReq, ProfileName: &handle, @@ -189,7 +188,8 @@ func doExtendSession(ctx context.Context, cmd *cobra.Command) error { client := proto.NewDaemonServiceClient(conn) - req := &proto.RequestExtendAuthSessionRequest{} + // the CLI runs in the user's session, the daemon does not: tell it what we can see + req := &proto.RequestExtendAuthSessionRequest{HasGraphicalSession: util.HasGraphicalSession()} // Pre-fill the IdP login hint from the active profile so the user // doesn't have to retype their email. Best-effort: we still proceed // without a hint if the lookup fails. @@ -408,7 +408,7 @@ func foregroundGetTokenInfo(ctx context.Context, cmd *cobra.Command, config *pro hint = profileState.Email } - oAuthFlow, err := auth.NewOAuthFlow(ctx, config, isUnixRunningDesktop(), false, hint) + oAuthFlow, err := auth.NewOAuthFlow(ctx, config, util.HasGraphicalSession(), false, hint) if err != nil { return nil, err } @@ -458,14 +458,6 @@ func openURL(cmd *cobra.Command, verificationURIComplete, userCode string, noBro } } -// isUnixRunningDesktop checks if a Linux OS is running desktop environment -func isUnixRunningDesktop() bool { - if runtime.GOOS != "linux" && runtime.GOOS != "freebsd" { - return false - } - return os.Getenv("DESKTOP_SESSION") != "" || os.Getenv("XDG_CURRENT_DESKTOP") != "" -} - func setEnvAndFlags(cmd *cobra.Command) error { SetFlagsFromEnvVars(rootCmd) diff --git a/client/cmd/logout.go b/client/cmd/logout.go index 1a5281acb..dcd7b5075 100644 --- a/client/cmd/logout.go +++ b/client/cmd/logout.go @@ -46,7 +46,7 @@ var logoutCmd = &cobra.Command{ } if _, err := daemonClient.Logout(ctx, req); err != nil { - return fmt.Errorf("deregister: %v", err) + return daemonCallError("deregister", err) } cmd.Println("Deregistered successfully") diff --git a/client/cmd/root.go b/client/cmd/root.go index f1ef32717..ccad78942 100644 --- a/client/cmd/root.go +++ b/client/cmd/root.go @@ -20,8 +20,8 @@ import ( "github.com/spf13/cobra" "github.com/spf13/pflag" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/anonymize" daddr "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/profilemanager" ) @@ -70,6 +70,7 @@ var ( autoConnectDisabled bool extraIFaceBlackList []string anonymizeFlag bool + anonymizeLevelFlag string dnsRouteInterval time.Duration // lazyConnEnabled is the parse target for the deprecated --enable-lazy-connection // flag. The flag is inert; the value is no longer read (use NB_LAZY_CONN instead). @@ -91,6 +92,7 @@ var ( // Don't resolve for service commands — they create the socket, not connect to it. if !isServiceCmd(cmd) { daemonAddr = daddr.ResolveUnixDaemonAddr(daemonAddr) + daemonAddr = daddr.ResolveDaemonAddr(daemonAddr) } return nil }, @@ -143,10 +145,10 @@ func init() { defaultDaemonAddr := "unix:///var/run/netbird.sock" if runtime.GOOS == "windows" { - defaultDaemonAddr = "tcp://127.0.0.1:41731" + defaultDaemonAddr = daddr.WindowsPipeAddr } - rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp]://[path|host:port]") + rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp|npipe]://[path|host:port|name]") rootCmd.PersistentFlags().StringVarP(&managementURL, "management-url", "m", "", fmt.Sprintf("Management Service URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultManagementURL)) rootCmd.PersistentFlags().StringVar(&adminURL, "admin-url", "", fmt.Sprintf("Admin Panel URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultAdminURL)) rootCmd.PersistentFlags().StringVarP(&logLevel, "log-level", "l", "info", "sets NetBird log level") @@ -156,7 +158,8 @@ func init() { rootCmd.MarkFlagsMutuallyExclusive("setup-key", "setup-key-file") rootCmd.PersistentFlags().StringVar(&preSharedKey, preSharedKeyFlag, "", "Sets WireGuard PreSharedKey property. If set, then only peers that have the same key can communicate.") rootCmd.PersistentFlags().StringVarP(&hostName, "hostname", "n", "", "Sets a custom hostname for the device") - rootCmd.PersistentFlags().BoolVarP(&anonymizeFlag, "anonymize", "A", false, "anonymize IP addresses and non-netbird.io domains in logs and status output") + rootCmd.PersistentFlags().BoolVarP(&anonymizeFlag, "anonymize", "A", false, "anonymize public IP addresses, MAC addresses, and non-netbird.io domains in logs and status output; private, CGNAT, and link-local IP ranges are kept (see --anonymize-level strict)") + rootCmd.PersistentFlags().StringVar(&anonymizeLevelFlag, "anonymize-level", "", "anonymization level: \"default\" or \"strict\"; strict also anonymizes private, CGNAT, and link-local IP ranges, peer names, and WireGuard public keys. Setting this flag implies --anonymize") rootCmd.PersistentFlags().StringVarP(&configPath, "config", "c", profilemanager.DefaultConfigPath, "Overrides the default profile file location") rootCmd.AddCommand(upCmd) @@ -269,12 +272,10 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e ctx, cancel := context.WithTimeout(ctx, time.Second*10) defer cancel() - return grpc.DialContext( - ctx, - strings.TrimPrefix(addr, "tcp://"), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpc.WithBlock(), - ) + target, opts := daddr.DialTarget(addr) + opts = append(opts, grpc.WithBlock()) + + return grpc.DialContext(ctx, target, opts...) } // WithBackOff execute function in backoff cycle. @@ -295,6 +296,19 @@ var CLIBackOffSettings = &backoff.ExponentialBackOff{ Clock: backoff.SystemClock, } +// effectiveAnonymize resolves the --anonymize and --anonymize-level flags: +// setting a level implies anonymization, and an invalid level is rejected. +func effectiveAnonymize() (bool, anonymize.Level, error) { + if anonymizeLevelFlag == "" { + return anonymizeFlag, anonymize.LevelDefault, nil + } + level := anonymize.ParseLevel(anonymizeLevelFlag) + if !strings.EqualFold(anonymizeLevelFlag, level.String()) { + return false, anonymize.LevelDefault, fmt.Errorf("invalid anonymize level %q: use %q or %q", anonymizeLevelFlag, anonymize.LevelDefault.String(), anonymize.LevelStrict.String()) + } + return true, level, nil +} + func getSetupKey() (string, error) { if setupKeyPath != "" && setupKey == "" { return getSetupKeyFromFile(setupKeyPath) diff --git a/client/cmd/service.go b/client/cmd/service.go index b0a56c71a..7410d60ea 100644 --- a/client/cmd/service.go +++ b/client/cmd/service.go @@ -33,10 +33,15 @@ var ( ) type program struct { - ctx context.Context - cancel context.CancelFunc - serv *grpc.Server - jsonServ *http.Server + ctx context.Context + cancel context.CancelFunc + serv *grpc.Server + jsonServ *http.Server + // jsonClient is the gateway's own connection to the daemon. It is held so + // shutting the gateway down also closes it: nothing else references it once + // the handlers are registered, so its transport goroutines would otherwise + // outlive the server. + jsonClient *grpc.ClientConn jsonServMu sync.Mutex serverInstance *server.Server serverInstanceMu sync.Mutex diff --git a/client/cmd/service_controller.go b/client/cmd/service_controller.go index 5ef13a0a6..b187a7b87 100644 --- a/client/cmd/service_controller.go +++ b/client/cmd/service_controller.go @@ -5,6 +5,7 @@ package cmd import ( "context" "fmt" + "runtime" "time" "github.com/kardianos/service" @@ -13,6 +14,8 @@ import ( "github.com/spf13/cobra" "google.golang.org/grpc" + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/server" "github.com/netbirdio/netbird/client/system" @@ -26,6 +29,31 @@ func validateJSONSocketFlags() error { return nil } +// daemonServerOptions installs the transport credentials that expose each +// caller's kernel-authenticated identity to the handlers, which is what lets +// the daemon require root/administrator for privileged operations. +// +// The handshake exchanges no bytes, so older CLI and UI binaries still +// interoperate. Callers on a TCP socket carry no identity at all: the daemon +// keeps serving them, and the privileged operations deny them, so a warning is +// logged to make the loss of functionality visible. +func daemonServerOptions(network string) []grpc.ServerOption { + if network == "tcp" { + log.Warnf("daemon is listening on TCP (%s): callers carry no verifiable identity over TCP, "+ + "so privileged operations (SSH root login, SSH auth, enabling the SSH server, management URL changes, "+ + "deregistration) will be denied. Use a unix socket, or npipe:// on Windows", daemonAddr) + return nil + } + + creds := ipcauth.NewTransportCredentials() //nolint:staticcheck + if creds == nil { //nolint:staticcheck // nil only on platforms without a peer-identity primitive + log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS) + return nil + } + + return []grpc.ServerOption{grpc.Creds(creds)} +} + func (p *program) Start(svc service.Service) error { // Start should not block. Do the actual work async. log.Info("starting NetBird service") //nolint @@ -37,68 +65,106 @@ func (p *program) Start(svc service.Service) error { // Collect static system and platform information system.UpdateStaticInfoAsync() - // in any case, even if configuration does not exists we run daemon to serve CLI gRPC API. - p.serv = grpc.NewServer() - - daemonListener, err := listenOnAddress(daemonAddr) - if err != nil { - return fmt.Errorf("listen daemon interface: %w", err) + // A daemon installed before named-pipe support has the loopback TCP address + // persisted. Move it to the named pipe so an upgraded daemon can identify + // its callers instead of silently serving an unauthenticated socket. + if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok { + log.Infof("daemon address %q predates named-pipe support, listening on %q so callers can be identified", daemonAddr, migrated) + daemonAddr = migrated } - var jsonListener *socketListener - if enableJSONSocket { - jsonListener, err = listenOnAddress(jsonSocket) - if err != nil { - _ = daemonListener.Close() - return fmt.Errorf("listen daemon JSON interface: %w", err) - } - } else { - removeStaleUnixSocketForAddress(jsonSocket) + network, _, err := parseListenAddress(daemonAddr) + if err != nil { + return fmt.Errorf("parse daemon address: %w", err) + } + + // in any case, even if configuration does not exists we run daemon to serve CLI gRPC API. + p.serv = grpc.NewServer(daemonServerOptions(network)...) + + daemonListener, jsonListener, err := listenDaemonSockets() + if err != nil { + return err } go func() { - defer daemonListener.Close() - if jsonListener != nil { - defer jsonListener.Close() - } - - if err := daemonListener.chmodUnixSocket("daemon"); err != nil { - log.Error(err) - return - } - if jsonListener != nil { - if err := jsonListener.chmodUnixSocket("daemon JSON"); err != nil { - log.Error(err) - return - } - } - - serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled) - if err := serverInstance.Start(); err != nil { - log.Fatalf("failed to start daemon: %v", err) - } - proto.RegisterDaemonServiceServer(p.serv, serverInstance) - - p.serverInstanceMu.Lock() - p.serverInstance = serverInstance - p.serverInstanceMu.Unlock() - - if jsonListener != nil { - if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil { - log.Fatalf("failed to start daemon JSON server: %v", err) - } - } else { - log.Debug("daemon JSON socket disabled") - } - - log.Printf("started daemon server: %v", daemonListener.address) - if err := p.serv.Serve(daemonListener.Listener); err != nil { - log.Errorf("failed to serve daemon requests: %v", err) + // Fatal here rather than inside serve, so serve's deferred listener + // closes run before the process exits. + if err := p.serve(daemonListener, jsonListener); err != nil { + log.Fatalf("failed to %v", err) } }() return nil } +// listenDaemonSockets opens the daemon control socket and, when it is enabled, the +// JSON gateway socket. The control socket is closed again if the second one fails, +// so a failed start leaves nothing listening. The returned JSON listener is nil +// when the socket is disabled. +func listenDaemonSockets() (*socketListener, *socketListener, error) { + daemonListener, err := listenOnAddress(daemonAddr) + if err != nil { + return nil, nil, fmt.Errorf("listen daemon interface: %w", err) + } + + if !enableJSONSocket { + removeStaleUnixSocketForAddress(jsonSocket) + return daemonListener, nil, nil + } + + jsonListener, err := listenOnAddress(jsonSocket) + if err != nil { + if cerr := daemonListener.Close(); cerr != nil { + log.Debugf("close daemon listener: %v", cerr) + } + return nil, nil, fmt.Errorf("listen daemon JSON interface: %w", err) + } + + return daemonListener, jsonListener, nil +} + +// serve brings up the daemon server on an already-open control socket and blocks +// until it stops. jsonListener is nil when the JSON socket is disabled. A returned +// error means the daemon cannot run at all and the caller is expected to exit; the +// failures it recovers from on its own are logged here. +func (p *program) serve(daemonListener, jsonListener *socketListener) error { + defer daemonListener.Close() + if jsonListener != nil { + defer jsonListener.Close() + } + + // chmodUnixSocket is a no-op for a nil listener and for a non-unix one. + if err := daemonListener.chmodUnixSocket("daemon"); err != nil { + log.Error(err) + return nil + } + if err := jsonListener.chmodUnixSocket("daemon JSON"); err != nil { + log.Error(err) + return nil + } + + serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled) + if err := serverInstance.Start(); err != nil { + return fmt.Errorf("start daemon: %w", err) + } + proto.RegisterDaemonServiceServer(p.serv, serverInstance) + + p.serverInstanceMu.Lock() + p.serverInstance = serverInstance + p.serverInstanceMu.Unlock() + + if jsonListener == nil { + log.Debug("daemon JSON socket disabled") + } else if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil { + return fmt.Errorf("start daemon JSON server: %w", err) + } + + log.Printf("started daemon server: %v", daemonListener.address) + if err := p.serv.Serve(daemonListener.Listener); err != nil { + log.Errorf("failed to serve daemon requests: %v", err) + } + return nil +} + func (p *program) Stop(srv service.Service) error { p.serverInstanceMu.Lock() if p.serverInstance != nil { @@ -113,8 +179,13 @@ func (p *program) Stop(srv service.Service) error { p.cancel() p.jsonServMu.Lock() - jsonServ := p.jsonServ + jsonServ, jsonClient := p.jsonServ, p.jsonClient p.jsonServMu.Unlock() + if jsonClient != nil { + if err := jsonClient.Close(); err != nil { + log.Debugf("close daemon JSON gateway client: %v", err) + } + } if jsonServ != nil { shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 2*time.Second) if err := jsonServ.Shutdown(shutdownCtx); err != nil { diff --git a/client/cmd/service_json_gateway.go b/client/cmd/service_json_gateway.go index 29c1a6456..b6864f338 100644 --- a/client/cmd/service_json_gateway.go +++ b/client/cmd/service_json_gateway.go @@ -5,27 +5,123 @@ package cmd import ( "context" "errors" + "fmt" "net" "net/http" - "strings" + "sync" "time" "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" log "github.com/sirupsen/logrus" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" ) -func grpcGatewayEndpoint(addr string) string { - return strings.TrimPrefix(addr, "tcp://") +// jsonPeerIdentity is the context key under which the connecting HTTP client's +// identity is stashed for the lifetime of its connection. +type jsonPeerIdentity struct{} + +// jsonPeerIdentityValue pairs the identity with whether it could be read at +// all, so an unreadable identity is forwarded as "unknown" rather than omitted. +type jsonPeerIdentityValue struct { + id ipcauth.Identity + known bool +} + +// jsonConnContext reads the identity of the client connecting to the JSON +// socket and stashes it on the connection's context. The gateway re-dials the +// daemon in-process, so the daemon would otherwise see every JSON request as +// coming from the daemon itself. +func jsonConnContext(ctx context.Context, c net.Conn) context.Context { + value := jsonPeerIdentityValue{} + id, err := ipcauth.ConnIdentity(c) + if err != nil { + log.Warnf("json gateway: cannot read HTTP client identity, privileged operations will be denied for this connection: %v", err) + } else { + value.id = id + value.known = true + } + return context.WithValue(ctx, jsonPeerIdentity{}, value) +} + +// forwardIdentity stamps the HTTP client's identity onto every call the gateway +// makes to the daemon. +// +// It is an interceptor on the gateway's client connection rather than a +// runtime.WithMetadata annotator because grpc-gateway skips annotators when no +// request header maps to metadata, which an HTTP/1.0 request with no Host header +// over a unix socket achieves. The daemon would then receive no marker, see its own +// identity as the transport peer, and authorize the request as the daemon itself. +// An interceptor runs for every RPC whatever the request looked like. +func forwardIdentity(ctx context.Context) context.Context { + value, ok := ctx.Value(jsonPeerIdentity{}).(jsonPeerIdentityValue) + if !ok { + // No ConnContext ran for this request, so forward an unknown identity: + // the daemon must not mistake its own identity for the client's. + return ipcauth.WithForwardedIdentity(ctx, ipcauth.Identity{}, false) + } + return ipcauth.WithForwardedIdentity(ctx, value.id, value.known) +} + +func forwardIdentityUnary(ctx context.Context, method string, req, reply any, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error { + return invoker(forwardIdentity(ctx), method, req, reply, cc, opts...) +} + +func forwardIdentityStream(ctx context.Context, desc *grpc.StreamDesc, cc *grpc.ClientConn, method string, streamer grpc.Streamer, opts ...grpc.CallOption) (grpc.ClientStream, error) { + return streamer(forwardIdentity(ctx), desc, cc, method, opts...) +} + +// reservedHeaderWarning limits the dropped-header warning to the first occurrence. +var reservedHeaderWarning sync.Once + +// jsonIncomingHeaderMatcher keeps an HTTP client from supplying the metadata the +// gateway uses to forward its identity. grpc-gateway turns "Grpc-Metadata-" +// headers into gRPC metadata and joins them ahead of what its annotators add, so +// without this filter a JSON client could send its own x-netbird-fwd-uid and the +// daemon would authorize that instead of the client's real identity. +func jsonIncomingHeaderMatcher(key string) (string, bool) { + mapped, ok := runtime.DefaultHeaderMatcher(key) + if !ok { + return "", false + } + if ipcauth.IsReservedForwardKey(mapped) { + // Warn once: any client can send these on every request, so warning each + // time hands it a way to fill the log. The rest are debug-level. + reservedHeaderWarning.Do(func() { + log.Warnf("json gateway: dropping reserved header %q from a request: only the gateway may set the caller's identity", key) + }) + log.Debugf("json gateway: dropping reserved header %q", key) + return "", false + } + return mapped, true } func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint string) error { - mux := runtime.NewServeMux() - opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())} - if err := proto.RegisterDaemonServiceHandlerFromEndpoint(p.ctx, mux, grpcGatewayEndpoint(daemonEndpoint), opts); err != nil { + if jsonListener.network == "tcp" { + log.Warnf("daemon JSON socket is listening on TCP (%s): callers carry no verifiable identity over TCP, "+ + "so privileged operations will be denied for JSON clients", jsonListener.address) + } + + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + + // grpc.NewClient does not connect until the first request, so registering + // the handler here cannot block daemon startup. + target, opts := daemonaddr.DialTarget(daemonEndpoint) + opts = append(opts, + grpc.WithChainUnaryInterceptor(forwardIdentityUnary), + grpc.WithChainStreamInterceptor(forwardIdentityStream), + ) + conn, err := grpc.NewClient(target, opts...) + if err != nil { + return fmt.Errorf("create daemon client for JSON gateway: %w", err) + } + if err := proto.RegisterDaemonServiceHandler(p.ctx, mux, conn); err != nil { + if cerr := conn.Close(); cerr != nil { + log.Debugf("close daemon client after failed JSON gateway registration: %v", cerr) + } return err } @@ -35,10 +131,12 @@ func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint BaseContext: func(net.Listener) context.Context { return p.ctx }, + ConnContext: jsonConnContext, } p.jsonServMu.Lock() p.jsonServ = jsonServer + p.jsonClient = conn p.jsonServMu.Unlock() go func() { diff --git a/client/cmd/service_json_gateway_test.go b/client/cmd/service_json_gateway_test.go new file mode 100644 index 000000000..dfeef1c46 --- /dev/null +++ b/client/cmd/service_json_gateway_test.go @@ -0,0 +1,261 @@ +//go:build !windows && !ios && !android + +package cmd + +import ( + "context" + "net" + "net/http" + "path/filepath" + "testing" + "time" + + "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/peer" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// The JSON gateway runs inside the daemon and re-dials it locally, so every JSON +// request reaches a handler with the daemon's own identity as the transport peer. +// The gateway therefore forwards its HTTP client's identity as metadata, and the +// daemon authorizes that instead of itself. These tests drive the real wiring +// (jsonConnContext, forwardIdentity, jsonIncomingHeaderMatcher) and check the +// identity a handler would end up authorizing. + +// daemonSideCtx is what a handler sees for a gateway-relayed call. The transport +// peer must be this process's own identity: the gateway is the daemon, so the two +// cannot differ, and hardcoding root here instead would describe a state that +// never occurs. +func daemonSideCtx(t *testing.T, md metadata.MD) context.Context { + t.Helper() + self, err := ipcauth.CurrentProcessIdentity() + if err != nil { + t.Skipf("cannot read this process's identity: %v", err) + } + ctx := peer.NewContext(context.Background(), &peer.Peer{ + AuthInfo: ipcauth.AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: self, + }, + }) + return metadata.NewIncomingContext(ctx, md) +} + +// gatewayMetadata reproduces what the daemon receives for a JSON request: the +// mux annotates the context from the request's headers, then the interceptor on the +// gateway's client connection stamps the caller's identity. The order matters, +// since the interceptor must win over anything a header put there. +func gatewayMetadata(t *testing.T, req *http.Request, ctx context.Context) metadata.MD { + t.Helper() + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + annotated, err := runtime.AnnotateContext(ctx, mux, req, + "/daemon.DaemonService/SetConfig", + runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig")) + if err != nil { + t.Fatalf("annotate: %v", err) + } + + md, ok := metadata.FromOutgoingContext(forwardIdentity(annotated)) + if !ok { + t.Fatal("the interceptor produced no metadata") + } + return md +} + +// clientCtx is the connection context jsonConnContext would have produced for an +// HTTP client whose identity the gateway could read. +func clientCtx(id ipcauth.Identity, known bool) context.Context { + return context.WithValue(context.Background(), jsonPeerIdentity{}, + jsonPeerIdentityValue{id: id, known: known}) +} + +// An HTTP client must not be able to name its own identity. grpc-gateway turns +// Grpc-Metadata- headers into gRPC metadata, so without the header filter and +// the interceptor overwriting the reserved keys, this request would authorize as +// uid 0. +func TestJSONGateway_ForgedIdentityHeaderIsDropped(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Uid", "0") + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Gid", "0") + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd", "1") + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Sid", "S-1-5-18") + + caller := ipcauth.Identity{UID: 31000, GID: 31000} + md := gatewayMetadata(t, req, clientCtx(caller, true)) + + id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)) + if !ok { + t.Fatal("the forwarded identity should be usable") + } + if id.IsPrivileged() { + t.Errorf("forged header was believed: authorized as %v", id) + } + if id.UID != caller.UID { + t.Errorf("authorized as uid %d, want the real client %d", id.UID, caller.UID) + } +} + +// A request with no headers at all (HTTP/1.0 needs no Host, and a unix socket +// yields no host:port) makes grpc-gateway produce no metadata whatsoever and skip +// its annotators: "if len(pairs) == 0 { return ctx, nil, nil }" in +// runtime/context.go. That is why the identity is stamped by an interceptor +// instead. This is the case that previously reached the gate as the daemon itself. +func TestJSONGateway_HeaderlessRequestIsStillMarkedForwarded(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + req.Header = http.Header{} + req.Host = "" + + caller := ipcauth.Identity{UID: 31000, GID: 31000} + ctx := clientCtx(caller, true) + + // Pin the skip path itself: if grpc-gateway ever produced a pair here, this + // test would still pass below while no longer covering what it was written for. + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + annotated, err := runtime.AnnotateContext(ctx, mux, req, + "/daemon.DaemonService/SetConfig", + runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig")) + if err != nil { + t.Fatalf("annotate: %v", err) + } + if md, ok := metadata.FromOutgoingContext(annotated); ok { + t.Fatalf("grpc-gateway produced metadata %v for a headerless request; "+ + "this test no longer covers the annotator-skip path", md) + } + + md := gatewayMetadata(t, req, ctx) + + id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)) + if !ok { + t.Fatal("the forwarded identity should be usable") + } + if id.UID != caller.UID || id.IsPrivileged() { + t.Errorf("authorized as %v, want the real client uid %d", id, caller.UID) + } +} + +// When the gateway cannot read its client's identity (a TCP JSON socket, say) it +// forwards the marker alone. The daemon must then report "unidentified" so the +// privileged operations refuse, rather than falling back to the gateway's own +// identity. +func TestJSONGateway_UnreadableClientIdentityIsUnidentified(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + + md := gatewayMetadata(t, req, clientCtx(ipcauth.Identity{}, false)) + + if id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)); ok { + t.Errorf("a request with no client identity was authorized as %v", id) + } +} + +// A request that never passed through jsonConnContext (no stashed identity) must +// also come out unidentified rather than as the daemon. +func TestJSONGateway_MissingConnContextIsUnidentified(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + + md := gatewayMetadata(t, req, context.Background()) + + if id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)); ok { + t.Errorf("a request with no connection context was authorized as %v", id) + } +} + +// End to end over a real unix socket: the gateway reads the connecting client's +// identity from the socket itself, so a client cannot present anything else. +func TestJSONGateway_IdentityComesFromTheSocket(t *testing.T) { + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + + type observed struct { + md metadata.MD + } + seen := make(chan observed, 1) + + srv := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ctx, err := runtime.AnnotateContext(r.Context(), mux, r, + "/daemon.DaemonService/SetConfig", + runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig")) + if err != nil { + t.Errorf("annotate: %v", err) + return + } + md, _ := metadata.FromOutgoingContext(forwardIdentity(ctx)) + seen <- observed{md: md} + w.WriteHeader(http.StatusOK) + }), + ReadHeaderTimeout: 5 * time.Second, + ConnContext: jsonConnContext, + } + + sock := filepath.Join(t.TempDir(), "http.sock") + ln, err := net.Listen("unix", sock) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := srv.Close(); err != nil { + t.Logf("close server: %v", err) + } + }) + go func() { + if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed { + t.Logf("serve: %v", err) + } + }() + + conn, err := net.Dial("unix", sock) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := conn.Close(); err != nil { + t.Logf("close conn: %v", err) + } + }) + + // Forge the identity headers on the wire as well. + request := "POST /daemon.DaemonService/SetConfig HTTP/1.1\r\n" + + "Host: localhost\r\n" + + "Grpc-Metadata-X-Netbird-Fwd: 1\r\n" + + "Grpc-Metadata-X-Netbird-Fwd-Uid: 0\r\n" + + "Content-Length: 0\r\n\r\n" + if _, err := conn.Write([]byte(request)); err != nil { + t.Fatal(err) + } + + select { + case got := <-seen: + self, err := ipcauth.CurrentProcessIdentity() + if err != nil { + t.Skipf("cannot read this process's identity: %v", err) + } + // The socket peer is this test process, so that is the identity the + // gateway must forward, not the uid 0 the request asked for. + if uids := got.md.Get("x-netbird-fwd-uid"); len(uids) != 1 { + t.Fatalf("x-netbird-fwd-uid = %v, want exactly the gateway's own value", uids) + } + id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, got.md)) + if !ok { + t.Fatal("the forwarded identity should be usable") + } + if id.UID != self.UID { + t.Errorf("authorized as uid %d, want the socket peer %d", id.UID, self.UID) + } + case <-time.After(5 * time.Second): + t.Fatal("the gateway never handled the request") + } +} diff --git a/client/cmd/service_params.go b/client/cmd/service_params.go index f25087a69..750b22ae6 100644 --- a/client/cmd/service_params.go +++ b/client/cmd/service_params.go @@ -13,6 +13,7 @@ import ( "github.com/spf13/cobra" "github.com/netbirdio/netbird/client/configs" + "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/util" ) @@ -125,6 +126,13 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) { if !rootCmd.PersistentFlags().Changed("daemon-addr") && params.DaemonAddr != "" { daemonAddr = params.DaemonAddr + // An install that predates named-pipe support has the loopback TCP + // address saved. Callers carry no identity over TCP, so move it to the + // pipe instead of restoring a socket the daemon cannot authorize on. + if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok { + cmd.Printf("Moving the saved daemon address from %s to %s so the daemon can identify its callers\n", daemonAddr, migrated) + daemonAddr = migrated + } } if !serviceCmd.PersistentFlags().Changed("json-socket") && params.JSONSocket != "" { diff --git a/client/cmd/service_pipe_other.go b/client/cmd/service_pipe_other.go new file mode 100644 index 000000000..c7cc72469 --- /dev/null +++ b/client/cmd/service_pipe_other.go @@ -0,0 +1,14 @@ +//go:build !windows + +package cmd + +import ( + "fmt" + "net" +) + +// listenNamedPipe is Windows-only: no other platform serves the daemon on a +// named pipe. +func listenNamedPipe(string) (net.Listener, string, error) { + return nil, "", fmt.Errorf("named pipes are only supported on Windows") +} diff --git a/client/cmd/service_pipe_windows.go b/client/cmd/service_pipe_windows.go new file mode 100644 index 000000000..b6e860f51 --- /dev/null +++ b/client/cmd/service_pipe_windows.go @@ -0,0 +1,41 @@ +//go:build windows + +package cmd + +import ( + "errors" + "fmt" + "net" + + "github.com/Microsoft/go-winio" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// listenNamedPipe creates the daemon control pipe and reports the path it ended +// up on. The security descriptor lets any local caller connect, as a Unix socket +// at 0666 does, and the privileged operations are authorized separately from the +// caller's token. +// +// The protected name comes first so that an unprivileged process cannot take the +// name before the service does. Creating it requires being an administrator or +// LocalSystem, so a daemon an ordinary user runs themselves, as in netstack mode, +// falls back to the plain name; clients try both and check who serves them. +func listenNamedPipe(name string) (net.Listener, string, error) { + var errs []error + for _, path := range daemonaddr.PipePaths(name) { + listener, err := winio.ListenPipe(path, &winio.PipeConfig{ + SecurityDescriptor: ipcauth.DefaultPipeSDDL(), + }) + if err != nil { + log.Debugf("not serving the daemon on %s: %v", path, err) + errs = append(errs, fmt.Errorf("%s: %w", path, err)) + continue + } + return listener, path, nil + } + + return nil, "", errors.Join(errs...) +} diff --git a/client/cmd/service_socket.go b/client/cmd/service_socket.go index f825a4062..bf3122f7c 100644 --- a/client/cmd/service_socket.go +++ b/client/cmd/service_socket.go @@ -26,6 +26,14 @@ func listenOnAddress(addr string) (*socketListener, error) { return nil, err } + if network == "npipe" { + listener, path, err := listenNamedPipe(address) //nolint:staticcheck + if err != nil { //nolint:staticcheck // always errors on non-Windows builds + return nil, err + } + return &socketListener{Listener: listener, network: network, address: path}, nil + } + if network == "unix" { removeStaleUnixSocket(address) } @@ -41,11 +49,11 @@ func listenOnAddress(addr string) (*socketListener, error) { func parseListenAddress(addr string) (string, string, error) { network, address, ok := strings.Cut(addr, "://") if !ok || network == "" || address == "" { - return "", "", fmt.Errorf("address must be in [unix|tcp]://[path|host:port] format: %q", addr) + return "", "", fmt.Errorf("address must be in [unix|tcp|npipe]://[path|host:port|name] format: %q", addr) } switch network { - case "unix", "tcp": + case "unix", "tcp", "npipe": return network, address, nil default: return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network) diff --git a/client/cmd/status.go b/client/cmd/status.go index c4057ed82..f2e5bcc66 100644 --- a/client/cmd/status.go +++ b/client/cmd/status.go @@ -121,8 +121,14 @@ func statusFunc(cmd *cobra.Command, args []string) error { sessionExpiresAt = ts.AsTime().UTC() } + anonymizeEnabled, anonymizeLevel, err := effectiveAnonymize() + if err != nil { + return err + } + var outputInformationHolder = nbstatus.ConvertToStatusOutputOverview(resp.GetFullStatus(), nbstatus.ConvertOptions{ - Anonymize: anonymizeFlag, + Anonymize: anonymizeEnabled, + AnonymizeLevel: anonymizeLevel, DaemonVersion: resp.GetDaemonVersion(), DaemonStatus: nbstatus.ParseDaemonStatus(status), StatusFilter: statusFilter, diff --git a/client/cmd/testutil_test.go b/client/cmd/testutil_test.go index 205327ef5..f40056f83 100644 --- a/client/cmd/testutil_test.go +++ b/client/cmd/testutil_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel" "google.golang.org/grpc" diff --git a/client/cmd/up.go b/client/cmd/up.go index 2d9731f26..9f4fa8c33 100644 --- a/client/cmd/up.go +++ b/client/cmd/up.go @@ -21,8 +21,8 @@ import ( "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/profilemanager" - "github.com/netbirdio/netbird/client/proto" nbnet "github.com/netbirdio/netbird/client/net" + "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/server" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/shared/management/domain" @@ -325,7 +325,7 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable { log.Warnf("setConfig method is not available in the daemon: %s", st.Message()) } else { - return fmt.Errorf("call service setConfig method: %v", err) + return daemonCallError("call service setConfig method", err) } } @@ -379,7 +379,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ } if loginErr != nil { - return fmt.Errorf("login failed: %v", loginErr) + return daemonCallError("login failed", loginErr) } if loginResp.NeedsSSOLogin { @@ -392,7 +392,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ ProfileName: &profileID, Username: &username, }); err != nil { - return fmt.Errorf("call service up method: %v", err) + return daemonCallError("call service up method", err) } return nil @@ -626,7 +626,7 @@ func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte NatExternalIPs: natExternalIPs, CleanNATExternalIPs: natExternalIPs != nil && len(natExternalIPs) == 0, CustomDNSAddress: customDNSAddressConverted, - IsUnixDesktopClient: isUnixRunningDesktop(), + IsUnixDesktopClient: util.HasGraphicalSession(), Hostname: hostName, ExtraIFaceBlacklist: extraIFaceBlackList, DnsLabels: dnsLabels, diff --git a/client/configs/configs.go b/client/configs/configs.go index 8f9c3ba28..a1ecf0feb 100644 --- a/client/configs/configs.go +++ b/client/configs/configs.go @@ -6,6 +6,11 @@ import ( "runtime" ) +// UILogFile is the file name the desktop UI writes its log to. It is defined +// here so the UI (writer), the daemon's RegisterUILog validation, and the debug +// bundle collector all share one definition. +const UILogFile = "gui-client.log" + var StateDir string func init() { diff --git a/client/embed/embed.go b/client/embed/embed.go index 99a6b8229..079e03c63 100644 --- a/client/embed/embed.go +++ b/client/embed/embed.go @@ -21,7 +21,7 @@ import ( "github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/profilemanager" - sshcommon "github.com/netbirdio/netbird/client/ssh" + nbssh "github.com/netbirdio/netbird/client/ssh" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/shared/management/domain" mgmProto "github.com/netbirdio/netbird/shared/management/proto" @@ -91,6 +91,13 @@ type Options struct { // when the embedded client must never act as a stepping stone into // the host's local network (e.g. the proxy's overlay peer). BlockLANAccess bool + // LazyConnectionEnabled is a tri-state local override for lazy connections, + // mirroring the NB_LAZY_CONN env var. Nil defers to the management feature + // flag; a set value overrides it in both directions. A short-lived client + // that reaches only a few known peers can set this to false, so its peers + // connect eagerly and the first request does not wait for the connection to + // be established. + LazyConnectionEnabled *bool // WireguardPort is the port for the tunnel interface. Use 0 for a random port. WireguardPort *int // MTU is the MTU for the tunnel interface. @@ -220,6 +227,15 @@ func New(opts Options) (*Client, error) { config.PrivateKey = opts.PrivateKey } + if opts.LazyConnectionEnabled != nil { + // Runtime-only override, read back through lazyconn.ParseState; a set value + // wins over the management feature flag in both directions. + config.LazyConnection = "off" + if *opts.LazyConnectionEnabled { + config.LazyConnection = "on" + } + } + if opts.Performance.PreallocatedBuffersPerPool != nil { wgdevice.SetPreallocatedBuffersPerPool(*opts.Performance.PreallocatedBuffersPerPool) } @@ -521,12 +537,7 @@ func (c *Client) VerifySSHHostKey(peerAddress string, key []byte) error { return err } - storedKey, found := engine.GetPeerSSHKey(peerAddress) - if !found { - return sshcommon.ErrPeerNotFound - } - - return sshcommon.VerifyHostKey(storedKey, key, peerAddress) + return nbssh.PeerKeyLookup(engine.GetPeerSSHKey).VerifySSHHostKey(peerAddress, key) } // SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool diff --git a/client/embed/embed_test.go b/client/embed/embed_test.go index a2f438975..27beb8934 100644 --- a/client/embed/embed_test.go +++ b/client/embed/embed_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/stretchr/testify/require" "google.golang.org/grpc" diff --git a/client/firewall/firewalld/firewalld.go b/client/firewall/firewalld/firewalld.go index 188ea61dd..38a4efdbc 100644 --- a/client/firewall/firewalld/firewalld.go +++ b/client/firewall/firewalld/firewalld.go @@ -2,8 +2,8 @@ // its wg interface into firewalld's "trusted" zone. This is required because // firewalld's nftables chains are created with NFT_CHAIN_OWNER on recent // versions, which returns EPERM to any other process that tries to insert -// rules into them. The workaround mirrors what Tailscale does: let firewalld -// itself add the accept rules to its own chains by trusting the interface. +// rules into them. Trusting the interface makes firewalld itself add the +// accept rules to its own chains instead. package firewalld // TrustedZone is the firewalld zone name used for interfaces whose traffic diff --git a/client/firewall/iptables/acl_linux.go b/client/firewall/iptables/acl_linux.go index 4b4cebf9c..89d1ebf7c 100644 --- a/client/firewall/iptables/acl_linux.go +++ b/client/firewall/iptables/acl_linux.go @@ -42,6 +42,7 @@ type aclManager struct { optionalEntries map[string][]entry ipsetStore *ipsetStore v6 bool + ipsetSupported bool stateManager *statemanager.Manager } @@ -60,6 +61,8 @@ func newAclManager(iptablesClient *iptables.IPTables, wgIface iFaceMapper) (*acl func (m *aclManager) init(stateManager *statemanager.Manager) error { m.stateManager = stateManager + m.ipsetSupported = m.probeIPSetSupport() + m.seedInitialEntries() m.seedInitialOptionalEntries() @@ -91,6 +94,12 @@ func (m *aclManager) AddPeerFiltering( if m.v6 && ipsetName != "" { ipsetName += "-v6" } + // When the kernel lacks the required ipset hash module, fall back to + // per-IP iptables rules (pre-0.68 behavior) so ACLs keep working instead + // of silently leaving the chain empty. + if ipsetName != "" && !m.ipsetSupported { + ipsetName = "" + } proto := protoForFamily(protocol, m.v6) specs := filterRuleSpecs(ip, proto, sPort, dPort, action, ipsetName) @@ -498,6 +507,40 @@ func transformIPsetName(ipsetName string, sPort, dPort *firewall.Port, action fi } } +// probeIPSetSupport checks whether the kernel can create the ipset type used for +// ACL rules. On kernels lacking the required ipset hash module, ipset creation +// fails (e.g. "invalid argument"), which would otherwise leave the ACL chain +// empty and silently drop all policy-permitted inbound traffic. When unsupported, +// the manager falls back to per-IP iptables rules. +func (m *aclManager) probeIPSetSupport() bool { + // Use a unique name so concurrent processes don't collide and we only ever + // destroy the set we created ourselves. ipset names are limited to 31 chars, + // so use a short random suffix. + probeName := "nb-probe-" + uuid.New().String()[:8] + + opts := ipset.CreateOptions{ + Replace: true, + } + if m.v6 { + opts.Family = ipset.FamilyIPV6 + } + + if err := ipset.Create(probeName, ipset.TypeHashNet, opts); err != nil { + log.Warnf("ipset is not available (failed to create probe set: %v); "+ + "falling back to per-IP iptables ACL rules. Ensure the kernel provides "+ + "the ipset hash:net module (ip_set_hash_net) for better performance with large rule sets", err) + return false + } + + defer func() { + if err := ipset.Destroy(probeName); err != nil { + log.Debugf("destroy ipset probe set %q: %v", probeName, err) + } + }() + + return true +} + func (m *aclManager) createIPSet(name string) error { opts := ipset.CreateOptions{ Replace: true, diff --git a/client/firewall/iptables/dnat_refcount_linux_test.go b/client/firewall/iptables/dnat_refcount_linux_test.go new file mode 100644 index 000000000..681bc0b99 --- /dev/null +++ b/client/firewall/iptables/dnat_refcount_linux_test.go @@ -0,0 +1,240 @@ +//go:build privileged + +package iptables + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + fw "github.com/netbirdio/netbird/client/firewall/manager" + "github.com/netbirdio/netbird/client/iface" + "github.com/netbirdio/netbird/client/iface/wgaddr" +) + +func iptRefcountIfaceV4() *iFaceMock { + return &iFaceMock{ + NameFunc: func() string { return "wt-refcount" }, + AddressFunc: func() wgaddr.Address { + return wgaddr.Address{ + IP: netip.MustParseAddr("10.20.0.1"), + Network: netip.MustParsePrefix("10.20.0.0/24"), + } + }, + } +} + +func iptRefcountIfaceDual() *iFaceMock { + return &iFaceMock{ + NameFunc: func() string { return "wt-refcount" }, + AddressFunc: func() wgaddr.Address { + return wgaddr.Address{ + IP: netip.MustParseAddr("10.20.0.1"), + Network: netip.MustParsePrefix("10.20.0.0/24"), + IPv6: netip.MustParseAddr("fd00::1"), + IPv6Net: netip.MustParsePrefix("fd00::/64"), + } + }, + } +} + +func newIptRefcountManager(t *testing.T, dual bool) *Manager { + t.Helper() + var ifMock *iFaceMock + if dual { + ifMock = iptRefcountIfaceDual() + } else { + ifMock = iptRefcountIfaceV4() + } + m, err := Create(ifMock, iface.DefaultMTU) + require.NoError(t, err, "create manager") + require.NoError(t, m.Init(nil), "init manager") + t.Cleanup(func() { + require.NoError(t, m.Close(nil), "close manager") + }) + return m +} + +func iptDnatV4(port uint16) fw.ForwardRule { + return fw.ForwardRule{ + Protocol: fw.ProtocolTCP, + DestinationPort: fw.Port{Values: []uint16{port}}, + TranslatedAddress: netip.MustParseAddr("10.20.0.2"), + TranslatedPort: fw.Port{Values: []uint16{80}}, + } +} + +func iptDnatV6(port uint16) fw.ForwardRule { + return fw.ForwardRule{ + Protocol: fw.ProtocolTCP, + DestinationPort: fw.Port{Values: []uint16{port}}, + TranslatedAddress: netip.MustParseAddr("fd00::2"), + TranslatedPort: fw.Port{Values: []uint16{80}}, + } +} + +// TestIptablesRouting_RepeatedEnableSingleReference verifies that EnableRouting +// (called on every network-map update) holds at most one reference per family +// and a single DisableRouting drops both back to zero. +func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + require.NoError(t, m.EnableRouting(), "first enable") + require.NoError(t, m.EnableRouting(), "second enable") + require.NoError(t, m.EnableRouting(), "third enable") + v4, v6 := state.Counts() + assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference") + assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference") + + require.NoError(t, m.DisableRouting(), "disable") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "single disable releases the v4 reference") + assert.Equal(t, 0, v6, "single disable releases the v6 reference") +} + +// TestIptablesRouting_DisableKeepsDNATReference verifies that an unpaired +// DisableRouting does not release references held by active DNAT rules. +func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(iptDnatV6(9095)) + require.NoError(t, err, "add v6 dnat") + + require.NoError(t, m.DisableRouting(), "unpaired disable") + _, v6 := state.Counts() + assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") + + require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat") + _, v6 = state.Counts() + assert.Equal(t, 0, v6, "delete releases the DNAT reference") +} + +// TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4. +func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) { + m := newIptRefcountManager(t, false) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(iptDnatV4(7081)) + require.NoError(t, err, "add v4 dnat 1") + v4, v6 := state.Counts() + assert.Equal(t, 1, v4, "v4 refcount after first add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") + + r2, err := m.AddDNATRule(iptDnatV4(7082)) + require.NoError(t, err, "add v4 dnat 2") + v4, v6 = state.Counts() + assert.Equal(t, 2, v4, "v4 refcount after second add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") + + require.NoError(t, m.DeleteDNATRule(r1)) + v4, v6 = state.Counts() + assert.Equal(t, 1, v4, "v4 refcount after first delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") + + require.NoError(t, m.DeleteDNATRule(r2)) + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "v4 refcount after second delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") +} + +// TestIptablesDNAT_RefcountBalancedV6 checks the v6 path increments v6 only and +// decrements back to zero. +func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) { + m := newIptRefcountManager(t, true) + require.NotNil(t, m.router6, "v6 router") + require.Same(t, m.router.ipFwdState, m.router6.ipFwdState, "shared state") + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(iptDnatV6(9081)) + require.NoError(t, err, "add v6 dnat 1") + v4, v6 := state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 1, v6, "v6 refcount after first add") + + r2, err := m.AddDNATRule(iptDnatV6(9082)) + require.NoError(t, err, "add v6 dnat 2") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 2, v6, "v6 refcount after second add") + + require.NoError(t, m.DeleteDNATRule(r1)) + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 1, v6, "v6 refcount after first delete") + + require.NoError(t, m.DeleteDNATRule(r2)) + v4, v6 = state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6, "v6 refcount after second delete") +} + +// TestIptablesDNAT_DuplicateAddNoLeak verifies the duplicate-rule path returns +// without bumping the refcount. +func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + rule := iptDnatV4(7083) + r1, err := m.AddDNATRule(rule) + require.NoError(t, err) + v4, _ := state.Counts() + assert.Equal(t, 1, v4) + + _, err = m.AddDNATRule(rule) + require.NoError(t, err, "duplicate add") + v4, _ = state.Counts() + assert.Equal(t, 1, v4, "duplicate add must not increment") + + require.NoError(t, m.DeleteDNATRule(r1)) + v4, _ = state.Counts() + assert.Equal(t, 0, v4, "single delete must drop to zero") +} + +// TestIptablesDNAT_DeleteMissingNoUnderflow verifies Delete on an unknown rule +// neither errors nor releases the refcount. +func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + phantom := iptDnatV4(7099) + require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4") + v4, v6 := state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6) + + phantom6 := iptDnatV6(9099) + require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6) + + r1, err := m.AddDNATRule(iptDnatV4(7100)) + require.NoError(t, err) + v4, _ = state.Counts() + assert.Equal(t, 1, v4, "real add still increments after phantom delete") + require.NoError(t, m.DeleteDNATRule(r1)) +} + +// TestIptablesDNAT_DoubleDeleteNoUnderflow verifies a second Delete on the same +// rule is a no-op. +func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { + m := newIptRefcountManager(t, true) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(iptDnatV6(9083)) + require.NoError(t, err) + _, v6 := state.Counts() + assert.Equal(t, 1, v6) + + require.NoError(t, m.DeleteDNATRule(r1), "first delete") + _, v6 = state.Counts() + assert.Equal(t, 0, v6) + + require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op") + _, v6 = state.Counts() + assert.Equal(t, 0, v6, "double delete must not underflow") +} diff --git a/client/firewall/iptables/manager_linux.go b/client/firewall/iptables/manager_linux.go index 696537dd8..aa052d933 100644 --- a/client/firewall/iptables/manager_linux.go +++ b/client/firewall/iptables/manager_linux.go @@ -89,7 +89,7 @@ func (m *Manager) createIPv6Components(wgIface iFaceMapper, mtu uint16) error { } // Share the same IP forwarding state with the v4 router, since - // EnableIPForwarding controls both v4 and v6 sysctls. + // Forwarding refcounter is per-family but shared between v4 and v6 routers. m.router6.ipFwdState = m.router.ipFwdState m.aclMgr6, err = newAclManager(ip6Client, wgIface) @@ -402,17 +402,12 @@ func (m *Manager) SetLogLevel(log.Level) { } func (m *Manager) EnableRouting() error { - if err := m.router.ipFwdState.RequestForwarding(); err != nil { - return fmt.Errorf("enable IP forwarding: %w", err) - } - return nil + // v6 only when the overlay actually has v6. + return m.router.ipFwdState.RequestRouting(m.router6 != nil) } func (m *Manager) DisableRouting() error { - if err := m.router.ipFwdState.ReleaseForwarding(); err != nil { - return fmt.Errorf("disable IP forwarding: %w", err) - } - return nil + return m.router.ipFwdState.ReleaseRouting() } // AddDNATRule adds a DNAT rule diff --git a/client/firewall/iptables/manager_linux_test.go b/client/firewall/iptables/manager_linux_test.go index 7b0989f6c..2c3c1a08e 100644 --- a/client/firewall/iptables/manager_linux_test.go +++ b/client/firewall/iptables/manager_linux_test.go @@ -291,3 +291,40 @@ func TestIptablesCreatePerformance(t *testing.T) { }) } } + +// TestIptablesACLIPSetFallback verifies that when the kernel lacks ipset support, +// the ACL manager falls back to per-IP iptables rules (-s ) instead of +// silently leaving the chain empty. See discussion #6125. +func TestIptablesACLIPSetFallback(t *testing.T) { + ipv4Client, err := iptables.NewWithProtocol(iptables.ProtocolIPv4) + require.NoError(t, err) + + // Use Create()/Init() so the router-owned chains (chainRTFWDIN/OUT) are + // created before the ACL manager's createDefaultChains() references them. + manager, err := Create(ifaceMock, iface.DefaultMTU) + require.NoError(t, err) + require.NoError(t, manager.Init(nil)) + + aclMgr := manager.aclMgr + // Simulate a kernel without the ipset hash module. + aclMgr.ipsetSupported = false + + defer func() { + require.NoError(t, manager.Close(nil)) + }() + + ip := netip.MustParseAddr("10.20.0.42") + port := &fw.Port{Values: []uint16{22}} + + rules, err := aclMgr.AddPeerFiltering(nil, ip.AsSlice(), "tcp", nil, port, fw.ActionAccept, "nb0000001") + require.NoError(t, err, "AddPeerFiltering should succeed via fallback") + require.NotEmpty(t, rules) + + rule := rules[0].(*Rule) + require.Empty(t, rule.ipsetName, "fallback rule must not reference an ipset") + require.Contains(t, strings.Join(rule.specs, " "), "-s 10.20.0.42", "fallback rule must match by source IP") + require.NotContains(t, strings.Join(rule.specs, " "), "--match-set", "fallback rule must not use ipset matching") + + // The rule must actually be present in the ACL chain (not silently dropped). + checkRuleSpecs(t, ipv4Client, rule.chain, true, rule.specs...) +} diff --git a/client/firewall/iptables/router_linux.go b/client/firewall/iptables/router_linux.go index 42d305f5c..01b18570c 100644 --- a/client/firewall/iptables/router_linux.go +++ b/client/firewall/iptables/router_linux.go @@ -102,7 +102,7 @@ func newRouter(iptablesClient *iptables.IPTables, wgIface iFaceMapper, mtu uint1 wgIface: wgIface, mtu: mtu, v6: iptablesClient.Proto() == iptables.ProtocolIPv6, - ipFwdState: ipfwdstate.NewIPForwardingState(), + ipFwdState: ipfwdstate.NewIPForwardingState(wgIface.Name()), } r.ipsetCounter = refcounter.New( @@ -770,10 +770,6 @@ func (r *router) updateState() { } func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { - if err := r.ipFwdState.RequestForwarding(); err != nil { - return nil, err - } - ruleKey := rule.ID() if _, exists := r.rules[ruleKey+dnatSuffix]; exists { return rule, nil @@ -840,18 +836,34 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { for key, ruleInfo := range rules { if err := r.iptablesClient.Append(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil { - if rollbackErr := r.rollbackRules(rules); rollbackErr != nil { - log.Errorf("rollback failed: %v", rollbackErr) - } + r.cleanupFailedDNATAdd(rules) return nil, fmt.Errorf("add rule %s: %w", key, err) } r.rules[key] = ruleInfo.rule } + if err := r.ipFwdState.RequestForwarding(r.v6); err != nil { + r.cleanupFailedDNATAdd(rules) + return nil, fmt.Errorf("enable forwarding: %w", err) + } + r.updateState() return rule, nil } +// cleanupFailedDNATAdd removes the bookkeeping written by a partially applied +// AddDNATRule before rolling back the kernel rules, so no entries remain that +// never got a forwarding refcount. rollbackRules re-adds entries it failed to +// remove from the kernel. +func (r *router) cleanupFailedDNATAdd(rules map[string]ruleInfo) { + for key := range rules { + delete(r.rules, key) + } + if err := r.rollbackRules(rules); err != nil { + log.Errorf("rollback failed: %v", err) + } +} + func (r *router) rollbackRules(rules map[string]ruleInfo) error { var merr *multierror.Error for key, ruleInfo := range rules { @@ -868,32 +880,47 @@ func (r *router) rollbackRules(rules map[string]ruleInfo) error { } func (r *router) DeleteDNATRule(rule firewall.Rule) error { - if err := r.ipFwdState.ReleaseForwarding(); err != nil { - log.Errorf("%v", err) - } - ruleKey := rule.ID() + _, hadDNAT := r.rules[ruleKey+dnatSuffix] + _, hadSNAT := r.rules[ruleKey+snatSuffix] + _, hadFWD := r.rules[ruleKey+fwdSuffix] + if !hadDNAT && !hadSNAT && !hadFWD { + return nil + } + var merr *multierror.Error if dnatRule, exists := r.rules[ruleKey+dnatSuffix]; exists { if err := r.iptablesClient.Delete(tableNat, chainRTRDR, dnatRule...); err != nil { merr = multierror.Append(merr, fmt.Errorf("delete DNAT rule: %w", err)) + } else { + delete(r.rules, ruleKey+dnatSuffix) } - delete(r.rules, ruleKey+dnatSuffix) } if snatRule, exists := r.rules[ruleKey+snatSuffix]; exists { if err := r.iptablesClient.Delete(tableNat, chainRTNAT, snatRule...); err != nil { merr = multierror.Append(merr, fmt.Errorf("delete SNAT rule: %w", err)) + } else { + delete(r.rules, ruleKey+snatSuffix) } - delete(r.rules, ruleKey+snatSuffix) } if fwdRule, exists := r.rules[ruleKey+fwdSuffix]; exists { if err := r.iptablesClient.Delete(tableFilter, chainRTFWDOUT, fwdRule...); err != nil { merr = multierror.Append(merr, fmt.Errorf("delete forward rule: %w", err)) + } else { + delete(r.rules, ruleKey+fwdSuffix) + } + } + + // Release the refcount only once all rules are gone from the kernel. On + // partial failure the failed entries stay in r.rules so a retry can remove + // them and release then. + if merr == nil { + if err := r.ipFwdState.ReleaseForwarding(r.v6); err != nil { + log.Errorf("%v", err) } - delete(r.rules, ruleKey+fwdSuffix) } r.updateState() diff --git a/client/firewall/nftables/dnat_refcount_linux_test.go b/client/firewall/nftables/dnat_refcount_linux_test.go new file mode 100644 index 000000000..86079676f --- /dev/null +++ b/client/firewall/nftables/dnat_refcount_linux_test.go @@ -0,0 +1,249 @@ +//go:build privileged + +package nftables + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + fw "github.com/netbirdio/netbird/client/firewall/manager" + "github.com/netbirdio/netbird/client/iface" + "github.com/netbirdio/netbird/client/iface/wgaddr" +) + +func nftRefcountIfaceV4() *iFaceMock { + return &iFaceMock{ + NameFunc: func() string { return "wt-refcount" }, + AddressFunc: func() wgaddr.Address { + return wgaddr.Address{ + IP: netip.MustParseAddr("100.96.0.1"), + Network: netip.MustParsePrefix("100.96.0.0/16"), + } + }, + } +} + +func nftRefcountIfaceDual() *iFaceMock { + return &iFaceMock{ + NameFunc: func() string { return "wt-refcount" }, + AddressFunc: func() wgaddr.Address { + return wgaddr.Address{ + IP: netip.MustParseAddr("100.96.0.1"), + Network: netip.MustParsePrefix("100.96.0.0/16"), + IPv6: netip.MustParseAddr("fd00::1"), + IPv6Net: netip.MustParsePrefix("fd00::/64"), + } + }, + } +} + +func newNftRefcountManager(t *testing.T, dual bool) *Manager { + t.Helper() + if check() != NFTABLES { + t.Skip("nftables not supported on this system") + } + var ifMock *iFaceMock + if dual { + ifMock = nftRefcountIfaceDual() + } else { + ifMock = nftRefcountIfaceV4() + } + m, err := Create(ifMock, iface.DefaultMTU) + require.NoError(t, err, "create manager") + require.NoError(t, m.Init(nil), "init manager") + t.Cleanup(func() { + require.NoError(t, m.Close(nil), "close manager") + }) + return m +} + +func dnatV4(port uint16) fw.ForwardRule { + return fw.ForwardRule{ + Protocol: fw.ProtocolTCP, + DestinationPort: fw.Port{Values: []uint16{port}}, + TranslatedAddress: netip.MustParseAddr("100.96.0.2"), + TranslatedPort: fw.Port{Values: []uint16{80}}, + } +} + +func dnatV6(port uint16) fw.ForwardRule { + return fw.ForwardRule{ + Protocol: fw.ProtocolTCP, + DestinationPort: fw.Port{Values: []uint16{port}}, + TranslatedAddress: netip.MustParseAddr("fd00::2"), + TranslatedPort: fw.Port{Values: []uint16{80}}, + } +} + +// TestNftablesDNAT_RefcountBalancedV4 verifies that Add/Delete pairs leave the +// v4 refcount at zero. +func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) { + m := newNftRefcountManager(t, false) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(dnatV4(8081)) + require.NoError(t, err, "add v4 dnat 1") + v4, v6 := state.Counts() + assert.Equal(t, 1, v4, "v4 refcount after first add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") + + r2, err := m.AddDNATRule(dnatV4(8082)) + require.NoError(t, err, "add v4 dnat 2") + v4, v6 = state.Counts() + assert.Equal(t, 2, v4, "v4 refcount after second add") + assert.Equal(t, 0, v6, "v6 refcount unchanged") + + require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat 1") + v4, v6 = state.Counts() + assert.Equal(t, 1, v4, "v4 refcount after first delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") + + require.NoError(t, m.DeleteDNATRule(r2), "delete v4 dnat 2") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "v4 refcount after second delete") + assert.Equal(t, 0, v6, "v6 refcount unchanged") +} + +// TestNftablesDNAT_RefcountBalancedV6 verifies the v6 path increments v6 only +// and decrements back to zero on Delete. +func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) { + m := newNftRefcountManager(t, true) + require.NotNil(t, m.router6, "v6 router") + require.Same(t, m.router.ipFwdState, m.router6.ipFwdState, "shared state") + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(dnatV6(9091)) + require.NoError(t, err, "add v6 dnat 1") + v4, v6 := state.Counts() + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 1, v6, "v6 refcount after first add") + + r2, err := m.AddDNATRule(dnatV6(9092)) + require.NoError(t, err, "add v6 dnat 2") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 2, v6, "v6 refcount after second add") + + require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat 1") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "v4 refcount unchanged") + assert.Equal(t, 1, v6, "v6 refcount after first delete") + + require.NoError(t, m.DeleteDNATRule(r2), "delete v6 dnat 2") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6, "v6 refcount after second delete") +} + +// TestNftablesDNAT_DuplicateAddNoLeak verifies that a duplicate Add (same +// ForwardRule) does not double-increment the refcount. +func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + rule := dnatV4(8083) + r1, err := m.AddDNATRule(rule) + require.NoError(t, err, "add v4 dnat") + v4, _ := state.Counts() + assert.Equal(t, 1, v4) + + // duplicate add: same rule ID, must be a no-op for the refcount. + _, err = m.AddDNATRule(rule) + require.NoError(t, err, "duplicate add") + v4, _ = state.Counts() + assert.Equal(t, 1, v4, "duplicate add must not increment") + + require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat") + v4, _ = state.Counts() + assert.Equal(t, 0, v4, "single delete must drop to zero") +} + +// TestNftablesDNAT_DeleteMissingNoUnderflow verifies deleting a rule that was +// never added does not underflow the refcount. +func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + // Construct a Rule reference for something never added. The router stores + // rules by ID(), and DeleteDNATRule looks them up in r.rules; a missing + // entry must be a no-op rather than calling Release. + phantom := dnatV4(8099) + require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4 dnat") + v4, v6 := state.Counts() + assert.Equal(t, 0, v4, "v4 refcount unaffected by missing delete") + assert.Equal(t, 0, v6, "v6 refcount unaffected") + + phantom6 := dnatV6(9099) + require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6 dnat") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4) + assert.Equal(t, 0, v6, "v6 refcount unaffected by missing delete") + + // And after a phantom delete, a real add still results in count=1. + r1, err := m.AddDNATRule(dnatV4(8100)) + require.NoError(t, err, "add v4 dnat after phantom delete") + v4, _ = state.Counts() + assert.Equal(t, 1, v4, "real add still increments after phantom delete") + require.NoError(t, m.DeleteDNATRule(r1)) +} + +// TestNftablesRouting_RepeatedEnableSingleReference verifies that EnableRouting +// (called on every network-map update) holds at most one reference per family +// and a single DisableRouting drops both back to zero. +func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + require.NoError(t, m.EnableRouting(), "first enable") + require.NoError(t, m.EnableRouting(), "second enable") + require.NoError(t, m.EnableRouting(), "third enable") + v4, v6 := state.Counts() + assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference") + assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference") + + require.NoError(t, m.DisableRouting(), "disable") + v4, v6 = state.Counts() + assert.Equal(t, 0, v4, "single disable releases the v4 reference") + assert.Equal(t, 0, v6, "single disable releases the v6 reference") +} + +// TestNftablesRouting_DisableKeepsDNATReference verifies that an unpaired +// DisableRouting does not release references held by active DNAT rules. +func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(dnatV6(9095)) + require.NoError(t, err, "add v6 dnat") + + require.NoError(t, m.DisableRouting(), "unpaired disable") + _, v6 := state.Counts() + assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") + + require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat") + _, v6 = state.Counts() + assert.Equal(t, 0, v6, "delete releases the DNAT reference") +} + +// TestNftablesDNAT_DoubleDeleteNoUnderflow verifies that deleting the same rule +// twice does not underflow the refcount (the second delete is a no-op). +func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { + m := newNftRefcountManager(t, true) + state := m.router.ipFwdState + + r1, err := m.AddDNATRule(dnatV6(9093)) + require.NoError(t, err) + _, v6 := state.Counts() + assert.Equal(t, 1, v6) + + require.NoError(t, m.DeleteDNATRule(r1), "first delete") + _, v6 = state.Counts() + assert.Equal(t, 0, v6) + + require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op") + _, v6 = state.Counts() + assert.Equal(t, 0, v6, "double delete must not underflow") +} diff --git a/client/firewall/nftables/legacy_rule_linux_test.go b/client/firewall/nftables/legacy_rule_linux_test.go new file mode 100644 index 000000000..dc2f1c7a0 --- /dev/null +++ b/client/firewall/nftables/legacy_rule_linux_test.go @@ -0,0 +1,60 @@ +package nftables + +import ( + "testing" + + "github.com/google/nftables/expr" + "github.com/stretchr/testify/require" +) + +func TestBuildLegacyRouteRuleExpressions(t *testing.T) { + sourcePayload := &expr.Payload{} + sourceCmp := &expr.Cmp{} + destinationPayload := &expr.Payload{} + destinationCmp := &expr.Cmp{} + nilSourceDestination := &expr.Payload{} + nilDestinationSource := &expr.Cmp{} + + tests := []struct { + name string + source []expr.Any + destination []expr.Any + matches []expr.Any + }{ + { + name: "both non-empty", + source: []expr.Any{sourcePayload, sourceCmp}, + destination: []expr.Any{destinationPayload, destinationCmp}, + matches: []expr.Any{sourcePayload, sourceCmp, destinationPayload, destinationCmp}, + }, + { + name: "nil source", + destination: []expr.Any{nilSourceDestination}, + matches: []expr.Any{nilSourceDestination}, + }, + { + name: "nil destination", + source: []expr.Any{nilDestinationSource}, + matches: []expr.Any{nilDestinationSource}, + }, + { + name: "both nil", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := buildLegacyRouteRuleExpressions(tt.source, tt.destination) + + require.Len(t, got, len(tt.matches)+2) + for i, match := range tt.matches { + require.Same(t, match, got[i]) + } + + require.IsType(t, &expr.Counter{}, got[len(tt.matches)]) + verdict, ok := got[len(tt.matches)+1].(*expr.Verdict) + require.True(t, ok) + require.Equal(t, expr.VerdictAccept, verdict.Kind) + }) + } +} diff --git a/client/firewall/nftables/manager_linux.go b/client/firewall/nftables/manager_linux.go index fdc7c2f3c..984b1c3ba 100644 --- a/client/firewall/nftables/manager_linux.go +++ b/client/firewall/nftables/manager_linux.go @@ -105,8 +105,8 @@ func (m *Manager) createIPv6Components(tableName string, wgIface iFaceMapper, mt return fmt.Errorf("create v6 router: %w", err) } - // Share the same IP forwarding state with the v4 router, since - // EnableIPForwarding controls both v4 and v6 sysctls. + // Share the per-family forwarding refcounter with the v4 router so a v4 + // rule and a v6 rule against the same state machine cooperate cleanly. m.router6.ipFwdState = m.router.ipFwdState m.aclManager6, err = newAclManager(workTable6, wgIface, chainNameRoutingFw) @@ -530,17 +530,12 @@ func (m *Manager) SetLogLevel(log.Level) { } func (m *Manager) EnableRouting() error { - if err := m.router.ipFwdState.RequestForwarding(); err != nil { - return fmt.Errorf("enable IP forwarding: %w", err) - } - return nil + // v6 only when the overlay actually has v6. + return m.router.ipFwdState.RequestRouting(m.router6 != nil) } func (m *Manager) DisableRouting() error { - if err := m.router.ipFwdState.ReleaseForwarding(); err != nil { - return fmt.Errorf("disable IP forwarding: %w", err) - } - return nil + return m.router.ipFwdState.ReleaseRouting() } // Flush rule/chain/set operations from the buffer diff --git a/client/firewall/nftables/router_linux.go b/client/firewall/nftables/router_linux.go index 4214455a9..d3e031c5f 100644 --- a/client/firewall/nftables/router_linux.go +++ b/client/firewall/nftables/router_linux.go @@ -93,7 +93,7 @@ func newRouter(workTable *nftables.Table, wgIface iFaceMapper, mtu uint16) (*rou rules: make(map[string]*nftables.Rule), af: familyForAddr(workTable.Family == nftables.TableFamilyIPv4), wgIface: wgIface, - ipFwdState: ipfwdstate.NewIPForwardingState(), + ipFwdState: ipfwdstate.NewIPForwardingState(wgIface.Name()), mtu: mtu, } @@ -953,6 +953,17 @@ func (r *router) addMSSClampingRules() error { return r.conn.Flush() } +func buildLegacyRouteRuleExpressions(sourceExp, destExp []expr.Any) []expr.Any { + exprs := make([]expr.Any, 0, len(sourceExp)+len(destExp)+2) + exprs = append(exprs, sourceExp...) + exprs = append(exprs, destExp...) + exprs = append(exprs, + &expr.Counter{}, + &expr.Verdict{Kind: expr.VerdictAccept}, + ) + return exprs +} + // addLegacyRouteRule adds a legacy routing rule for mgmt servers pre route acls func (r *router) addLegacyRouteRule(pair firewall.RouterPair) error { sourceExp, err := r.applyNetwork(pair.Source, nil, true) @@ -965,15 +976,7 @@ func (r *router) addLegacyRouteRule(pair firewall.RouterPair) error { return fmt.Errorf("apply destination: %w", err) } - exprs := []expr.Any{ - &expr.Counter{}, - &expr.Verdict{ - Kind: expr.VerdictAccept, - }, - } - - exprs = append(exprs, sourceExp...) - exprs = append(exprs, destExp...) + exprs := buildLegacyRouteRuleExpressions(sourceExp, destExp) ruleKey := firewall.GenKey(firewall.ForwardingFormat, pair) @@ -1550,10 +1553,6 @@ func (r *router) refreshRulesMap() error { } func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { - if err := r.ipFwdState.RequestForwarding(); err != nil { - return nil, err - } - ruleKey := rule.ID() if _, exists := r.rules[ruleKey+dnatSuffix]; exists { return rule, nil @@ -1564,7 +1563,18 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { return nil, fmt.Errorf("convert protocol to number: %w", err) } + // Request forwarding before queueing rules: addDnatRedirect/addDnatMasq + // buffer netlink messages on r.conn that the next caller's Flush would + // commit if we returned without flushing them ourselves. + v6 := r.af.tableFamily == nftables.TableFamilyIPv6 + if err := r.ipFwdState.RequestForwarding(v6); err != nil { + return nil, fmt.Errorf("enable forwarding: %w", err) + } + if err := r.addDnatRedirect(rule, protoNum, ruleKey); err != nil { + if rerr := r.ipFwdState.ReleaseForwarding(v6); rerr != nil { + log.Warnf("rollback forwarding refcount: %v", rerr) + } return nil, err } @@ -1576,6 +1586,11 @@ func (r *router) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { // TODO: find chains with drop policies and add rules there if err := r.conn.Flush(); err != nil { + if rerr := r.ipFwdState.ReleaseForwarding(v6); rerr != nil { + log.Warnf("rollback forwarding refcount: %v", rerr) + } + delete(r.rules, ruleKey+dnatSuffix) + delete(r.rules, ruleKey+snatSuffix) return nil, fmt.Errorf("flush rules: %w", err) } @@ -1778,16 +1793,18 @@ func (r *router) addDnatMasq(rule firewall.ForwardRule, protoNum uint8, ruleKey } func (r *router) DeleteDNATRule(rule firewall.Rule) error { - if err := r.ipFwdState.ReleaseForwarding(); err != nil { - log.Errorf("%v", err) - } - ruleKey := rule.ID() if err := r.refreshRulesMap(); err != nil { return fmt.Errorf(refreshRulesMapError, err) } + _, hadDNAT := r.rules[ruleKey+dnatSuffix] + _, hadSNAT := r.rules[ruleKey+snatSuffix] + if !hadDNAT && !hadSNAT { + return nil + } + var merr *multierror.Error var needsFlush bool @@ -1819,9 +1836,16 @@ func (r *router) DeleteDNATRule(rule firewall.Rule) error { } } + // Release the refcount only once the rules are gone from the kernel. On + // failure (including the refreshRulesMap error above) the rules and their + // map entries remain, keeping forwarding on until a retry removes them. if merr == nil { delete(r.rules, ruleKey+dnatSuffix) delete(r.rules, ruleKey+snatSuffix) + + if err := r.ipFwdState.ReleaseForwarding(r.af.tableFamily == nftables.TableFamilyIPv6); err != nil { + log.Errorf("%v", err) + } } return nberrors.FormatErrorOrNil(merr) diff --git a/client/firewall/uspfilter/filter_filter_test.go b/client/firewall/uspfilter/filter_filter_test.go index a64c83138..5ca8538be 100644 --- a/client/firewall/uspfilter/filter_filter_test.go +++ b/client/firewall/uspfilter/filter_filter_test.go @@ -5,7 +5,7 @@ import ( "net/netip" "testing" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/google/gopacket" "github.com/google/gopacket/layers" "github.com/stretchr/testify/require" diff --git a/client/firewall/uspfilter/filter_routeacl_test.go b/client/firewall/uspfilter/filter_routeacl_test.go index 449554d8b..b6397d09b 100644 --- a/client/firewall/uspfilter/filter_routeacl_test.go +++ b/client/firewall/uspfilter/filter_routeacl_test.go @@ -4,7 +4,7 @@ import ( "net/netip" "testing" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/google/gopacket/layers" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" diff --git a/client/grpc/dialer_generic.go b/client/grpc/dialer_generic.go index 479575996..737787223 100644 --- a/client/grpc/dialer_generic.go +++ b/client/grpc/dialer_generic.go @@ -16,28 +16,52 @@ import ( "google.golang.org/grpc" nbnet "github.com/netbirdio/netbird/client/net" + "github.com/netbirdio/netbird/client/netevents/sweep" ) +// Sweeper registers in-flight dials for the network change sweep. +type Sweeper interface { + StartDial(ctx context.Context) *sweep.Dial +} + func WithCustomDialer(_ bool, _ string) grpc.DialOption { + return grpc.WithContextDialer(dialContext) +} + +// WithSweeper dials like WithCustomDialer but registers connections and +// dials with the sweeper. Append it after WithCustomDialer: gRPC applies +// dial options in order, so the later context dialer wins. +func WithSweeper(sweeper Sweeper) grpc.DialOption { return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) { - if runtime.GOOS == "linux" { - currentUser, err := user.Current() - if err != nil { - return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err) - } + dial := sweeper.StartDial(ctx) + defer dial.Release() - // the custom dialer requires root permissions which are not required for use cases run as non-root - if currentUser.Uid != "0" { - log.Debug("Not running as root, using standard dialer") - dialer := &net.Dialer{} - return dialer.DialContext(ctx, "tcp", addr) - } - } - - conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr) + conn, err := dialContext(dial.Ctx(), addr) if err != nil { - return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err) + return nil, err } - return conn, nil + return dial.WrapConn(conn) }) } + +func dialContext(ctx context.Context, addr string) (net.Conn, error) { + if runtime.GOOS == "linux" { + currentUser, err := user.Current() + if err != nil { + return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err) + } + + // the custom dialer requires root permissions which are not required for use cases run as non-root + if currentUser.Uid != "0" { + log.Debug("Not running as root, using standard dialer") + dialer := &net.Dialer{} + return dialer.DialContext(ctx, "tcp", addr) + } + } + + conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr) + if err != nil { + return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err) + } + return conn, nil +} diff --git a/client/grpc/dialer_js.go b/client/grpc/dialer_js.go index b89ec3c21..4ff4ceb20 100644 --- a/client/grpc/dialer_js.go +++ b/client/grpc/dialer_js.go @@ -1,13 +1,26 @@ package grpc import ( + "context" + "google.golang.org/grpc" + "github.com/netbirdio/netbird/client/netevents/sweep" "github.com/netbirdio/netbird/util/wsproxy/client" ) +// Sweeper registers in-flight dials for the network change sweep. +type Sweeper interface { + StartDial(ctx context.Context) *sweep.Dial +} + // WithCustomDialer returns a gRPC dial option that uses WebSocket transport for WASM/JS environments. // The component parameter specifies the WebSocket proxy component path (e.g., "/management", "/signal"). func WithCustomDialer(tlsEnabled bool, component string) grpc.DialOption { return client.WithWebSocketDialer(tlsEnabled, component) } + +// WithSweeper is a no-op on WASM/JS: there is no network change signal. +func WithSweeper(_ Sweeper) grpc.DialOption { + return grpc.EmptyDialOption{} +} diff --git a/client/grpc/retry.go b/client/grpc/retry.go new file mode 100644 index 000000000..0bb6037bf --- /dev/null +++ b/client/grpc/retry.go @@ -0,0 +1,56 @@ +package grpc + +import ( + "context" + "errors" + "time" + + "github.com/cenkalti/backoff/v4" +) + +// ChangeWatcher exposes OS network availability transitions. +type ChangeWatcher interface { + Changed() <-chan struct{} +} + +// Retry mirrors backoff.Retry, but the sleep between attempts also wakes on +// OS network availability transitions: an operation cut down by a network +// change retries the moment the network settles instead of sleeping through +// the recovery. A nil watcher never fires, leaving plain backoff.Retry +// behavior. +func Retry(ctx context.Context, operation backoff.Operation, bo backoff.BackOff, watcher ChangeWatcher) error { + bo.Reset() + for { + err := operation() + if err == nil { + return nil + } + + var permanent *backoff.PermanentError + if errors.As(err, &permanent) { + return permanent.Err + } + + next := bo.NextBackOff() + if next == backoff.Stop { + if cerr := ctx.Err(); cerr != nil { + return cerr + } + return err + } + + var changed <-chan struct{} + if watcher != nil { + changed = watcher.Changed() + } + timer := time.NewTimer(next) + select { + case <-timer.C: + case <-changed: + timer.Stop() + case <-ctx.Done(): + timer.Stop() + return ctx.Err() + } + } +} diff --git a/client/grpc/retry_test.go b/client/grpc/retry_test.go new file mode 100644 index 000000000..266bb93e5 --- /dev/null +++ b/client/grpc/retry_test.go @@ -0,0 +1,91 @@ +package grpc + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/cenkalti/backoff/v4" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/netevents/netstate" +) + +func TestRetryWakesOnNetworkChange(t *testing.T) { + ns := netstate.New() + attempts := 0 + operation := func() error { + attempts++ + if attempts == 1 { + return errors.New("cut by network change") + } + return nil + } + + go func() { + time.Sleep(20 * time.Millisecond) + ns.Set(false) + }() + + start := time.Now() + err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Minute), ns) + + require.NoError(t, err) + assert.Equal(t, 2, attempts, "network change must cause one immediate retry") + assert.Less(t, time.Since(start), time.Second, "the transition must cut the minute-long sleep short") +} + +func TestRetryPermanentError(t *testing.T) { + sentinel := errors.New("permission denied") + operation := func() error { + return backoff.Permanent(sentinel) + } + + err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Millisecond), nil) + assert.ErrorIs(t, err, sentinel, "permanent errors must stop retries") +} + +func TestRetryNilNetState(t *testing.T) { + attempts := 0 + operation := func() error { + attempts++ + if attempts < 3 { + return errors.New("transient") + } + return nil + } + + err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Millisecond), nil) + require.NoError(t, err) + assert.Equal(t, 3, attempts, "nil network state must preserve timed retries") +} + +func TestRetryStops(t *testing.T) { + failure := errors.New("still failing") + operation := func() error { + return failure + } + + err := Retry(context.Background(), operation, &backoff.StopBackOff{}, nil) + assert.ErrorIs(t, err, failure, "stop backoff must return the operation error") +} + +func TestRetryCtxCancelDuringSleep(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + operation := func() error { + return errors.New("failing") + } + + go func() { + time.Sleep(20 * time.Millisecond) + cancel() + }() + + start := time.Now() + err := Retry(ctx, operation, backoff.NewConstantBackOff(time.Minute), netstate.New()) + + assert.ErrorIs(t, err, context.Canceled, "context cancellation must stop the retry loop") + assert.Less(t, time.Since(start), time.Second, "context cancellation must interrupt backoff sleep") +} diff --git a/client/iface/bind/ice_bind.go b/client/iface/bind/ice_bind.go index 156450c61..2d35b9c6f 100644 --- a/client/iface/bind/ice_bind.go +++ b/client/iface/bind/ice_bind.go @@ -22,6 +22,16 @@ import ( nbnet "github.com/netbirdio/netbird/client/net" ) +const ( + // wgMsgTypeHandshakeInitiation is the lowest WireGuard message type. + wgMsgTypeHandshakeInitiation uint32 = 1 + // wgMsgTypeTransport is the highest WireGuard message type. + wgMsgTypeTransport uint32 = 4 + // wgMinMsgSize is the smallest WireGuard message: transport data with an empty + // payload, which is what a keepalive is. + wgMinMsgSize = 32 +) + type receiverCreator struct { iceBind *ICEBind } @@ -216,8 +226,15 @@ func (s *ICEBind) createReceiverFn(pc wgConn.BatchReader, conn *net.UDPConn, rxO for i := 0; i < numMsgs; i++ { msg := &(*msgs)[i] - // todo: handle err - if ok, _ := s.filterOutStunMessages(msg.Buffers, msg.N, msg.Addr); ok { + if ok, err := s.filterOutStunMessages(msg.Buffers, msg.N, msg.Addr); ok { + if err != nil { + log.Debugf("failed to handle STUN packet from %s: %v", msg.Addr, err) + } + // WireGuard reuses sizes and eps across reads and only skips a slot + // whose size is below the minimum message size. Leaving a consumed + // slot untouched makes it process this buffer again under the + // previous packet's length and endpoint. + sizes[i] = 0 continue } sizes[i] = msg.N @@ -271,11 +288,16 @@ func (s *ICEBind) createOrUpdateMux() { func (s *ICEBind) filterOutStunMessages(buffers [][]byte, n int, addr net.Addr) (bool, error) { for i := range buffers { - if !stun.IsMessage(buffers[i]) { + if n > len(buffers[i]) { + continue + } + pkt := buffers[i][:n] + + if isWireGuardMsg(pkt) || !stun.IsMessage(pkt) { continue } - msg, err := s.parseSTUNMessage(buffers[i][:n]) + msg, err := s.parseSTUNMessage(pkt) if err != nil { buffers[i] = []byte{} return true, err @@ -347,18 +369,34 @@ func putMessages(msgs *[]ipv6.Message, msgsPool *sync.Pool) { msgsPool.Put(msgs) } -func isTransportPkg(buffers [][]byte, n int) bool { - // The first buffer should contain at least 4 bytes for type - if len(buffers[0]) < 4 { - return true +// isWireGuardMsg reports whether the packet carries a WireGuard message header: a +// little-endian uint32 message type in the range 1..4, which leaves the three bytes +// after the type byte zero, in a packet long enough to hold any WireGuard message. +// +// A well formed STUN message cannot take that shape. Its length field sits in the two +// bytes the type must leave zero, and for a message of at least wgMinMsgSize bytes that +// field holds at least 12, so the two framings do not overlap. The test has to be this +// tight because stun.IsMessage only looks at the magic cookie, which in a WireGuard +// message overlaps the receiver index: a session whose index happens to equal the cookie +// would otherwise have all of its inbound data misrouted to the STUN handler until the +// next rekey. +func isWireGuardMsg(pkt []byte) bool { + if len(pkt) < wgMinMsgSize { + return false } - // WireGuard packet type is a little-endian uint32 at start - packetType := binary.LittleEndian.Uint32(buffers[0][:4]) - - // Check if packetType matches known WireGuard message types - if packetType == 4 && n > 32 { - return true - } - return false + msgType := binary.LittleEndian.Uint32(pkt[:4]) + return msgType >= wgMsgTypeHandshakeInitiation && msgType <= wgMsgTypeTransport +} + +// isTransportPkg reports whether the packet is WireGuard transport data carrying a +// payload, which is what counts as peer activity. A keepalive holds no payload and is +// exactly wgMinMsgSize bytes. +func isTransportPkg(buffers [][]byte, n int) bool { + if n < 4 || n > len(buffers[0]) { + return false + } + + msgType := binary.LittleEndian.Uint32(buffers[0][:4]) + return msgType == wgMsgTypeTransport && n > wgMinMsgSize } diff --git a/client/iface/bind/stun_filter_test.go b/client/iface/bind/stun_filter_test.go new file mode 100644 index 000000000..0e118e0fd --- /dev/null +++ b/client/iface/bind/stun_filter_test.go @@ -0,0 +1,215 @@ +//go:build !js + +package bind + +import ( + "encoding/binary" + "net" + "testing" + "time" + + "github.com/pion/stun/v3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/net/ipv4" + wgConn "golang.zx2c4.com/wireguard/conn" +) + +// magicCookieBytes is the STUN magic cookie as it appears on the wire. In a +// WireGuard message the same offset holds the receiver (or sender) index, which is +// a random uint32, so a session can draw exactly this value. +var magicCookieBytes = []byte{0x21, 0x12, 0xA4, 0x42} + +const testBufSize = 1500 + +// wgMsg builds a WireGuard message of the given type and size, with the index field +// at bytes 4:8 set to index. +func wgMsg(msgType uint32, size int, index []byte) []byte { + pkt := make([]byte, size) + binary.LittleEndian.PutUint32(pkt[:4], msgType) + copy(pkt[4:8], index) + return pkt +} + +// intoBuffer copies pkt into a full-size receive buffer, the way the kernel read +// does, so tests see the same buffer/length split as the hot path. +func intoBuffer(pkt []byte) [][]byte { + buf := make([]byte, testBufSize) + copy(buf, pkt) + return [][]byte{buf} +} + +func TestFilterOutStunMessages_PassesWireGuardWithCookieShapedIndex(t *testing.T) { + tests := []struct { + name string + msgType uint32 + size int + }{ + {"transport data", wgMsgTypeTransport, 128}, + {"keepalive", wgMsgTypeTransport, wgMinMsgSize}, + {"handshake initiation", wgMsgTypeHandshakeInitiation, 148}, + {"handshake response", 2, 92}, + {"cookie reply", 3, 64}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + pkt := wgMsg(tc.msgType, tc.size, magicCookieBytes) + require.True(t, stun.IsMessage(pkt), "precondition: pion sees this as STUN") + + buffers := intoBuffer(pkt) + bind := &ICEBind{} + + filtered, err := bind.filterOutStunMessages(buffers, tc.size, &net.UDPAddr{}) + assert.NoError(t, err) + assert.False(t, filtered, "WireGuard message must be handed to WireGuard, not the STUN handler") + assert.Len(t, buffers[0], testBufSize, "buffer must be left intact for WireGuard") + }) + } +} + +func TestFilterOutStunMessages_FiltersRealSTUNMessage(t *testing.T) { + msg, err := stun.Build(stun.BindingRequest, stun.TransactionID, stun.Fingerprint) + require.NoError(t, err) + + buffers := intoBuffer(msg.Raw) + bind := &ICEBind{} + + filtered, err := bind.filterOutStunMessages(buffers, len(msg.Raw), &net.UDPAddr{}) + assert.NoError(t, err) + assert.True(t, filtered, "STUN message must be consumed by the STUN handler") + assert.Empty(t, buffers[0], "consumed buffer must be emptied so WireGuard does not see it") +} + +// TestIsWireGuardMsg_DisjointFromSTUN locks the invariant the filter relies on: a +// well formed STUN message long enough to be a WireGuard message always has a +// non-zero length field, so it cannot be mistaken for a WireGuard header. +func TestIsWireGuardMsg_DisjointFromSTUN(t *testing.T) { + types := []stun.MessageType{ + stun.BindingRequest, + stun.BindingSuccess, + stun.BindingError, + {Method: stun.MethodBinding, Class: stun.ClassIndication}, + } + + for _, msgType := range types { + // Long enough that the length guard is not what makes this pass. + msg, err := stun.Build(msgType, stun.TransactionID, + stun.NewUsername("remoteUfrag:localUfrag"), stun.Fingerprint) + require.NoError(t, err) + require.GreaterOrEqual(t, len(msg.Raw), wgMinMsgSize, "precondition: %s", msgType) + assert.False(t, isWireGuardMsg(msg.Raw), + "%s must not look like a WireGuard message", msgType) + } +} + +func TestIsWireGuardMsg(t *testing.T) { + tests := []struct { + name string + pkt []byte + want bool + }{ + {"transport data", wgMsg(wgMsgTypeTransport, 128, nil), true}, + {"handshake initiation", wgMsg(wgMsgTypeHandshakeInitiation, 148, nil), true}, + {"unknown type 5", wgMsg(5, 128, nil), false}, + {"type 0", wgMsg(0, 128, nil), false}, + {"non-zero reserved byte", []byte{0x04, 0x00, 0x01, 0x00}, false}, + {"too short", []byte{0x04, 0x00, 0x00}, false}, + {"empty", nil, false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, isWireGuardMsg(tc.pkt), "wrong classification for %s", tc.name) + }) + } +} + +// TestFilterOutStunMessages_IgnoresBytesBeyondPacket guards against classifying on +// buffer contents left over from an earlier, longer packet. +func TestFilterOutStunMessages_IgnoresBytesBeyondPacket(t *testing.T) { + buf := make([]byte, testBufSize) + copy(buf[4:8], magicCookieBytes) + buffers := [][]byte{buf} + bind := &ICEBind{} + + filtered, err := bind.filterOutStunMessages(buffers, 2, &net.UDPAddr{}) + assert.NoError(t, err) + assert.False(t, filtered, "a 2 byte packet must not be classified from stale buffer bytes") +} + +// TestReceiveFn_ClearsSizeOfConsumedPacket covers the accounting WireGuard relies +// on: sizes is reused across reads, so a slot whose packet was consumed as STUN must +// be reported as empty. Otherwise WireGuard reprocesses the same buffer under the +// previous packet's length, which for a WireGuard-shaped packet means it is handled +// twice. +func TestReceiveFn_ClearsSizeOfConsumedPacket(t *testing.T) { + conn := listenUDP(t, "udp4", "127.0.0.1:0") + defer conn.Close() + + recvFn := receiverCreator{setupICEBind(t)}.CreateReceiverFn( + ipv4.NewPacketConn(conn), conn, false, createMsgPool(), + ) + + msg, err := stun.Build(stun.BindingRequest, stun.TransactionID, stun.Fingerprint) + require.NoError(t, err) + + sender := listenUDP(t, "udp4", "127.0.0.1:0") + defer sender.Close() + _, err = sender.WriteTo(msg.Raw, conn.LocalAddr()) + require.NoError(t, err) + + require.NoError(t, conn.SetReadDeadline(time.Now().Add(3*time.Second))) + + bufs := [][]byte{make([]byte, 1500)} + // A leftover size from an earlier read, which is what makes the missing reset + // observable. + sizes := []int{148} + eps := make([]wgConn.Endpoint, 1) + + n, err := recvFn(bufs, sizes, eps) + require.NoError(t, err) + require.Equal(t, 1, n) + assert.Zero(t, sizes[0], "consumed STUN packet must not leave a size behind for WireGuard") +} + +func TestIsTransportPkg(t *testing.T) { + tests := []struct { + name string + pkt []byte + n int + want bool + }{ + {"transport data with payload", wgMsg(wgMsgTypeTransport, 128, nil), 128, true}, + {"keepalive", wgMsg(wgMsgTypeTransport, wgMinMsgSize, nil), wgMinMsgSize, false}, + {"handshake initiation", wgMsg(wgMsgTypeHandshakeInitiation, 148, nil), 148, false}, + {"stale type bytes beyond packet", wgMsg(wgMsgTypeTransport, 128, nil), 2, false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.want, isTransportPkg(intoBuffer(tc.pkt), tc.n), + "wrong activity classification for %s", tc.name) + }) + } +} + +// TestFilterOutStunMessages_ConsumesSTUNWithWireGuardShapedType covers the one STUN +// encoding whose leading bytes collide with a WireGuard message type: method 0x080 as a +// request encodes to 0x0200, so the type byte reads as a handshake response and the byte +// after it is zero. Only the length check keeps such a message out of WireGuard's hands. +// pion implements no method in that range, so this is a synthetic worst case rather than +// traffic ICE produces. +func TestFilterOutStunMessages_ConsumesSTUNWithWireGuardShapedType(t *testing.T) { + msg, err := stun.Build(stun.NewType(stun.Method(0x080), stun.ClassRequest), stun.TransactionID) + require.NoError(t, err) + require.Equal(t, []byte{0x02, 0x00, 0x00, 0x00}, msg.Raw[:4], + "precondition: the leading bytes read as a WireGuard message type") + + buffers := intoBuffer(msg.Raw) + bind := &ICEBind{} + + filtered, err := bind.filterOutStunMessages(buffers, len(msg.Raw), &net.UDPAddr{}) + assert.NoError(t, err) + assert.True(t, filtered, "STUN message must be consumed despite its WireGuard-shaped type") +} diff --git a/client/iface/device/device_filter_test.go b/client/iface/device/device_filter_test.go index 0d86c9323..a75ef90f9 100644 --- a/client/iface/device/device_filter_test.go +++ b/client/iface/device/device_filter_test.go @@ -4,7 +4,7 @@ import ( "net" "testing" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/google/gopacket" "github.com/google/gopacket/layers" diff --git a/client/iface/iface_test.go b/client/iface/iface_test.go index 8ff2bbb54..89c8cd16e 100644 --- a/client/iface/iface_test.go +++ b/client/iface/iface_test.go @@ -464,6 +464,8 @@ func Test_RemovePeer(t *testing.T) { } func Test_ConnectPeers(t *testing.T) { + t.Setenv("NB_DISABLE_EBPF_WG_PROXY", "true") + peer1ifaceName := fmt.Sprintf("utun%d", WgIntNumber+400) peer1wgIP := netip.MustParsePrefix("10.99.99.17/30") peer1Key, _ := wgtypes.GeneratePrivateKey() @@ -505,12 +507,8 @@ func Test_ConnectPeers(t *testing.T) { t.Fatal(err) } - localIP, err := getLocalIP() - if err != nil { - t.Fatal(err) - } - - peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer1wgPort)) + localIP1 := "127.0.0.1" + peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP1, peer1wgPort)) if err != nil { t.Fatal(err) } @@ -546,7 +544,8 @@ func Test_ConnectPeers(t *testing.T) { t.Fatal(err) } - peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer2wgPort)) + localIP2 := "127.0.0.1" + peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP2, peer2wgPort)) if err != nil { t.Fatal(err) } @@ -569,17 +568,17 @@ func Test_ConnectPeers(t *testing.T) { if err != nil { t.Fatal(err) } - // todo: investigate why in some tests execution we need 30s + // 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. timeout := 30 * time.Second timeoutChannel := time.After(timeout) + ticker := time.NewTicker(500 * time.Millisecond) + defer ticker.Stop() for { - select { - case <-timeoutChannel: - t.Fatalf("waiting for peer handshake timeout after %s", timeout.String()) - default: - } - peer, gpErr := getPeer(peer1ifaceName, peer2Key.PublicKey().String()) if gpErr != nil { t.Fatal(gpErr) @@ -588,6 +587,12 @@ func Test_ConnectPeers(t *testing.T) { t.Log("peers successfully handshake") break } + + select { + case <-timeoutChannel: + t.Fatalf("waiting for peer handshake timeout after %s", timeout.String()) + case <-ticker.C: + } } } @@ -615,28 +620,3 @@ func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) { } return wgtypes.Peer{}, fmt.Errorf("peer not found") } - -func getLocalIP() (string, error) { - // Get all interfaces - addrs, err := net.InterfaceAddrs() - if err != nil { - return "", err - } - - for _, addr := range addrs { - ipNet, ok := addr.(*net.IPNet) - if !ok { - continue - } - if ipNet.IP.IsLoopback() { - continue - } - - if ipNet.IP.To4() == nil { - continue - } - return ipNet.IP.String(), nil - } - - return "", fmt.Errorf("no local IP found") -} diff --git a/client/iface/mocks/filter.go b/client/iface/mocks/filter.go index 5ae98039c..ff3dd0c8a 100644 --- a/client/iface/mocks/filter.go +++ b/client/iface/mocks/filter.go @@ -8,7 +8,7 @@ import ( "net/netip" reflect "reflect" - gomock "github.com/golang/mock/gomock" + gomock "go.uber.org/mock/gomock" ) // MockPacketFilter is a mock of PacketFilter interface. diff --git a/client/iface/mocks/tun.go b/client/iface/mocks/tun.go index 677c82b0b..519ee6005 100644 --- a/client/iface/mocks/tun.go +++ b/client/iface/mocks/tun.go @@ -8,7 +8,7 @@ import ( os "os" reflect "reflect" - gomock "github.com/golang/mock/gomock" + gomock "go.uber.org/mock/gomock" tun "golang.zx2c4.com/wireguard/tun" ) diff --git a/client/iface/wgproxy/bind/proxy.go b/client/iface/wgproxy/bind/proxy.go index be690ed4f..fcaee15c7 100644 --- a/client/iface/wgproxy/bind/proxy.go +++ b/client/iface/wgproxy/bind/proxy.go @@ -53,15 +53,15 @@ func NewProxyBind(bind Bind, mtu uint16) *ProxyBind { return p } -// AddTurnConn adds a new connection to the bind. +// AddRelayedConn adds a new connection to the bind. // endpoint is the NetBird address of the remote peer. The SetEndpoint return with the address what will be used in the // WireGuard configuration. // // Parameters: // - ctx: Context is used for proxyToLocal to avoid unnecessary error messages // - nbAddr: The NetBird UDP address of the remote peer, it required to generate fake address -// - remoteConn: The established TURN connection to the remote peer -func (p *ProxyBind) AddTurnConn(ctx context.Context, nbAddr *net.UDPAddr, remoteConn net.Conn) error { +// - remoteConn: The established relayed connection to the remote peer +func (p *ProxyBind) AddRelayedConn(ctx context.Context, nbAddr *net.UDPAddr, remoteConn net.Conn) error { fakeNetIP, err := fakeAddress(nbAddr) if err != nil { return err diff --git a/client/iface/wgproxy/ebpf/proxy.go b/client/iface/wgproxy/ebpf/proxy.go index 1b1a8ce1c..91c741c0d 100644 --- a/client/iface/wgproxy/ebpf/proxy.go +++ b/client/iface/wgproxy/ebpf/proxy.go @@ -30,9 +30,9 @@ type WGEBPFProxy struct { proxyPort int mtu uint16 - ebpfManager ebpfMgr.Manager - turnConnStore map[uint16]net.Conn - turnConnMutex sync.Mutex + ebpfManager ebpfMgr.Manager + relayedConnStore map[uint16]net.Conn + relayedConnMutex sync.Mutex lastUsedPort uint16 rawConnIPv4 net.PacketConn @@ -50,7 +50,7 @@ func NewWGEBPFProxy(wgPort int, mtu uint16) *WGEBPFProxy { localWGListenPort: wgPort, mtu: mtu, ebpfManager: ebpf.GetEbpfManagerInstance(), - turnConnStore: make(map[uint16]net.Conn), + relayedConnStore: make(map[uint16]net.Conn), } return wgProxy } @@ -110,14 +110,14 @@ func (p *WGEBPFProxy) Listen() error { return nil } -// AddTurnConn add new turn connection for the proxy -func (p *WGEBPFProxy) AddTurnConn(turnConn net.Conn) (*net.UDPAddr, error) { - wgEndpointPort, err := p.storeTurnConn(turnConn) +// 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("turn conn added to wg proxy store: %s, endpoint port: :%d", turnConn.RemoteAddr(), wgEndpointPort) + log.Infof("relayed conn added to wg proxy store: %s, endpoint port: :%d", relayedConn.RemoteAddr(), wgEndpointPort) wgEndpoint := &net.UDPAddr{ IP: net.ParseIP(loopbackAddr), @@ -186,48 +186,48 @@ func (p *WGEBPFProxy) readAndForwardPacket(buf []byte) error { return fmt.Errorf("failed to read UDP packet from WG: %w", err) } - p.turnConnMutex.Lock() - conn, ok := p.turnConnStore[uint16(addr.Port)] - p.turnConnMutex.Unlock() + p.relayedConnMutex.Lock() + conn, ok := p.relayedConnStore[uint16(addr.Port)] + p.relayedConnMutex.Unlock() if !ok { if p.ctx.Err() == nil { - log.Debugf("turn conn not found by port because conn already has been closed: %d", addr.Port) + 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("failed to forward local WG packet (%d) to remote turn conn: %w", addr.Port, err) + return fmt.Errorf("forward local WG packet (%d) to remote relayed conn: %w", addr.Port, err) } return nil } -func (p *WGEBPFProxy) storeTurnConn(turnConn net.Conn) (uint16, error) { - p.turnConnMutex.Lock() - defer p.turnConnMutex.Unlock() +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.turnConnStore[np] = turnConn + p.relayedConnStore[np] = relayedConn return np, nil } -func (p *WGEBPFProxy) removeTurnConn(turnConnID uint16) { - p.turnConnMutex.Lock() - defer p.turnConnMutex.Unlock() +func (p *WGEBPFProxy) removeRelayedConn(relayedConnID uint16) { + p.relayedConnMutex.Lock() + defer p.relayedConnMutex.Unlock() - _, ok := p.turnConnStore[turnConnID] + _, ok := p.relayedConnStore[relayedConnID] if ok { - log.Debugf("remove turn conn from store by port: %d", turnConnID) + log.Debugf("remove relayed conn from store by port: %d", relayedConnID) } - delete(p.turnConnStore, turnConnID) + delete(p.relayedConnStore, relayedConnID) } func (p *WGEBPFProxy) nextFreePort() (uint16, error) { - if len(p.turnConnStore) == 65535 { - return 0, fmt.Errorf("reached maximum turn connection numbers") + if len(p.relayedConnStore) == 65535 { + return 0, fmt.Errorf("reached maximum relayed connection numbers") } generatePort: if p.lastUsedPort == 65535 { @@ -236,7 +236,7 @@ generatePort: p.lastUsedPort++ } - if _, ok := p.turnConnStore[p.lastUsedPort]; ok { + 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 index 3ec4f0eba..228c06c9b 100644 --- a/client/iface/wgproxy/ebpf/proxy_test.go +++ b/client/iface/wgproxy/ebpf/proxy_test.go @@ -9,32 +9,32 @@ import ( func TestWGEBPFProxy_connStore(t *testing.T) { wgProxy := NewWGEBPFProxy(1, 1280) - p, _ := wgProxy.storeTurnConn(nil) + 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.storeTurnConn(nil) + p, _ = wgProxy.storeRelayedConn(nil) } if p != uint16(numOfConns)+1 { t.Errorf("invalid last used port: %d, expected: %d", p, numOfConns+1) } - if len(wgProxy.turnConnStore) != numOfConns+1 { - t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.turnConnStore), 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.storeTurnConn(nil) + _, _ = wgProxy.storeRelayedConn(nil) wgProxy.lastUsedPort = 65535 - p, _ := wgProxy.storeTurnConn(nil) + p, _ := wgProxy.storeRelayedConn(nil) - if len(wgProxy.turnConnStore) != 2 { - t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.turnConnStore), 2) + if len(wgProxy.relayedConnStore) != 2 { + t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), 2) } if p != 2 { @@ -46,11 +46,11 @@ func TestWGEBPFProxy_portCalculation_maxConn(t *testing.T) { wgProxy := NewWGEBPFProxy(1, 1280) for i := 0; i < 65535; i++ { - _, _ = wgProxy.storeTurnConn(nil) + _, _ = wgProxy.storeRelayedConn(nil) } - _, err := wgProxy.storeTurnConn(nil) + _, err := wgProxy.storeRelayedConn(nil) if err == nil { - t.Errorf("invalid turn conn store calculation") + t.Errorf("invalid relayed conn store calculation") } } diff --git a/client/iface/wgproxy/ebpf/wrapper.go b/client/iface/wgproxy/ebpf/wrapper.go index a6156a661..f75e21aa6 100644 --- a/client/iface/wgproxy/ebpf/wrapper.go +++ b/client/iface/wgproxy/ebpf/wrapper.go @@ -121,10 +121,10 @@ func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper { } } -func (p *ProxyWrapper) AddTurnConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { - addr, err := p.wgeBPFProxy.AddTurnConn(remoteConn) +func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { + addr, err := p.wgeBPFProxy.AddRelayedConn(remoteConn) if err != nil { - return fmt.Errorf("add turn conn: %w", err) + return fmt.Errorf("add relayed conn: %w", err) } headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr) @@ -252,7 +252,7 @@ func (p *ProxyWrapper) CloseConn() error { } func (p *ProxyWrapper) proxyToLocal(ctx context.Context) { - defer p.wgeBPFProxy.removeTurnConn(uint16(p.wgRelayedEndpointAddr.Port)) + defer p.wgeBPFProxy.removeRelayedConn(uint16(p.wgRelayedEndpointAddr.Port)) buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead) for { @@ -273,7 +273,7 @@ func (p *ProxyWrapper) proxyToLocal(ctx context.Context) { if ctx.Err() != nil { return } - log.Errorf("failed to write out turn pkg to local conn: %v", err) + log.Errorf("failed to write out relayed pkg to local conn: %v", err) } } } @@ -286,7 +286,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 turn conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err) + log.Errorf("failed to read from relayed conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err) } return 0, err } diff --git a/client/iface/wgproxy/proxy.go b/client/iface/wgproxy/proxy.go index 40346bc15..b0033bffa 100644 --- a/client/iface/wgproxy/proxy.go +++ b/client/iface/wgproxy/proxy.go @@ -7,7 +7,7 @@ import ( // Proxy is a transfer layer between the relayed connection and the WireGuard type Proxy interface { - AddTurnConn(ctx context.Context, endpoint *net.UDPAddr, remoteConn net.Conn) error + AddRelayedConn(ctx context.Context, endpoint *net.UDPAddr, remoteConn net.Conn) error EndpointAddr() *net.UDPAddr // EndpointAddr returns the address of the WireGuard peer endpoint Work() // Work start or resume the proxy Pause() // Pause to forward the packages from remote connection to WireGuard. The opposite way still works. diff --git a/client/iface/wgproxy/proxy_test.go b/client/iface/wgproxy/proxy_test.go index 1aeab66b7..d86cdbe80 100644 --- a/client/iface/wgproxy/proxy_test.go +++ b/client/iface/wgproxy/proxy_test.go @@ -95,7 +95,7 @@ func TestProxyCloseByRemoteConn(t *testing.T) { t.Run(tt.name, func(t *testing.T) { addr, _ := net.ResolveUDPAddr("udp", "100.108.135.221:51892") relayedConn := newMockConn() - err := tt.proxy.AddTurnConn(ctx, addr, relayedConn) + err := tt.proxy.AddRelayedConn(ctx, addr, relayedConn) if err != nil { t.Errorf("error: %v", err) } @@ -157,7 +157,7 @@ func redirectTraffic(t *testing.T, proxy Proxy, wgPort int, endPointAddr *net.UD _ = relayedServer.Close() }() - if err := proxy.AddTurnConn(context.Background(), endPointAddr, relayedConn); err != nil { + if err := proxy.AddRelayedConn(context.Background(), endPointAddr, relayedConn); err != nil { t.Errorf("error: %v", err) } defer func() { diff --git a/client/iface/wgproxy/redirect_test.go b/client/iface/wgproxy/redirect_test.go index 135970838..f0d59cc64 100644 --- a/client/iface/wgproxy/redirect_test.go +++ b/client/iface/wgproxy/redirect_test.go @@ -119,9 +119,9 @@ func testRedirectAs(t *testing.T, proxy Proxy, wgPort int, nbAddr, p2pEndpoint * } defer relayConn.Close() - // Add TURN connection to proxy - if err := proxy.AddTurnConn(ctx, nbAddr, relayConn); err != nil { - t.Fatalf("failed to add TURN connection: %v", err) + // Add relayed connection to proxy + if err := proxy.AddRelayedConn(ctx, nbAddr, relayConn); err != nil { + t.Fatalf("failed to add relayed connection: %v", err) } defer func() { if err := proxy.CloseConn(); err != nil { @@ -304,8 +304,8 @@ func TestRedirectAs_Multiple_Switches(t *testing.T) { Port: 38746, } - if err := proxy.AddTurnConn(ctx, nbAddr, relayConn); err != nil { - t.Fatalf("failed to add TURN connection: %v", err) + if err := proxy.AddRelayedConn(ctx, nbAddr, relayConn); err != nil { + t.Fatalf("failed to add relayed connection: %v", err) } defer func() { if err := proxy.CloseConn(); err != nil { diff --git a/client/iface/wgproxy/udp/proxy.go b/client/iface/wgproxy/udp/proxy.go index 783843aba..a0895c8c7 100644 --- a/client/iface/wgproxy/udp/proxy.go +++ b/client/iface/wgproxy/udp/proxy.go @@ -51,12 +51,12 @@ func NewWGUDPProxy(wgPort int, mtu uint16) *WGUDPProxy { return p } -// AddTurnConn +// AddRelayedConn dials the local WireGuard port and stores the relayed connection. // The provided Context must be non-nil. If the context expires before // the connection is complete, an error is returned. Once successfully // connected, any expiration of the context will not affect the // connection. -func (p *WGUDPProxy) AddTurnConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { +func (p *WGUDPProxy) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { dialer := net.Dialer{} localConn, err := dialer.DialContext(ctx, "udp", fmt.Sprintf(":%d", p.localWGListenPort)) if err != nil { diff --git a/client/installer.nsis b/client/installer.nsis index 71699071b..eb2d7d5bd 100644 --- a/client/installer.nsis +++ b/client/installer.nsis @@ -22,8 +22,6 @@ !define UI_REG_APP_PATH "Software\Microsoft\Windows\CurrentVersion\App Paths\${UI_APP_EXE}" !define UI_UNINSTALL_PATH "Software\Microsoft\Windows\CurrentVersion\Uninstall\${UI_APP_NAME}" -!define AUTOSTART_REG_KEY "Software\Microsoft\Windows\CurrentVersion\Run" - !define NETBIRD_DATA_DIR "$COMMONPROGRAMDATA\Netbird" Unicode True @@ -228,13 +226,6 @@ WriteRegStr ${REG_ROOT} "${UNINSTALL_PATH}" "Publisher" "${COMP_NAME}" WriteRegStr ${REG_ROOT} "${UI_REG_APP_PATH}" "" "$INSTDIR\${UI_APP_EXE}" -; Autostart is owned by the UI's per-user setting (HKCU\...\Run via Wails), -; not the installer. Drop the machine-wide entry older installers wrote so the -; toggle is the single source of truth. HKCU is left untouched -- it may hold -; the user's own toggle state, which must survive upgrades. -DetailPrint "Removing installer-managed autostart registry entry if present..." -DeleteRegValue HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}" - EnVar::SetHKLM EnVar::AddValueEx "path" "$INSTDIR" @@ -299,15 +290,6 @@ ExecWait '"$INSTDIR\${MAIN_APP_EXE}" service uninstall' DetailPrint "Terminating Netbird UI process..." ExecWait `taskkill /im ${UI_APP_EXE}.exe /f` -; Remove autostart registry entries -DetailPrint "Removing autostart registry entries if they exist..." -; Legacy machine-wide entry written by older installers. -DeleteRegValue HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}" -; Per-user entry the UI toggle writes via Wails (value name is the lowercase -; app-name slug). Uninstall removes the app, so drop it too. -DeleteRegValue HKCU "${AUTOSTART_REG_KEY}" "${APP_NAME}" -DeleteRegValue HKCU "${AUTOSTART_REG_KEY}" "netbird" - ; Handle data deletion based on checkbox DetailPrint "Checking if user requested data deletion..." ${If} $DeleteDataEnabled == "1" diff --git a/client/internal/acl/manager.go b/client/internal/acl/manager.go index d9b179457..cbd9c5ab1 100644 --- a/client/internal/acl/manager.go +++ b/client/internal/acl/manager.go @@ -116,11 +116,11 @@ func (d *DefaultManager) ApplyFiltering(networkMap *mgmProto.NetworkMap, dnsRout // firewall state, so an identical hash means an identical resulting ruleset. func (d *DefaultManager) firewallConfigHash(networkMap *mgmProto.NetworkMap, dnsRouteFeatureFlag bool) (uint64, error) { return hashstructure.Hash(struct { - PeerRules []*mgmProto.FirewallRule - PeerRulesIsEmpty bool - RouteRules []*mgmProto.RouteFirewallRule - RouteRulesIsEmpty bool - DNSRouteFeatureFlag bool + PeerRules []*mgmProto.FirewallRule + PeerRulesIsEmpty bool + RouteRules []*mgmProto.RouteFirewallRule + RouteRulesIsEmpty bool + DNSRouteFeatureFlag bool }{ PeerRules: networkMap.GetFirewallRules(), PeerRulesIsEmpty: networkMap.GetFirewallRulesIsEmpty(), @@ -144,13 +144,13 @@ func (d *DefaultManager) applyPeerACLs(networkMap *mgmProto.NetworkMap) { log.Warn("this peer is connected to a NetBird Management service with an older version. Allowing all traffic from connected peers") rules = append(rules, &mgmProto.FirewallRule{ - PeerIP: "0.0.0.0", + PeerIP: "0.0.0.0", //nolint:staticcheck Direction: mgmProto.RuleDirection_IN, Action: mgmProto.RuleAction_ACCEPT, Protocol: mgmProto.RuleProtocol_ALL, }, &mgmProto.FirewallRule{ - PeerIP: "0.0.0.0", + PeerIP: "0.0.0.0", //nolint:staticcheck Direction: mgmProto.RuleDirection_OUT, Action: mgmProto.RuleAction_ACCEPT, Protocol: mgmProto.RuleProtocol_ALL, @@ -407,7 +407,6 @@ func (d *DefaultManager) getRuleGroupingSelector(rule *mgmProto.FirewallRule) st return fmt.Sprintf("%v:%v:%v:%s:%v", strconv.Itoa(int(rule.Direction)), rule.Action, rule.Protocol, rule.Port, rule.PortInfo) } - // extractRuleIP extracts the peer IP from a firewall rule. // If sourcePrefixes is populated (new management), decode the first entry and use its address. // Otherwise fall back to the deprecated PeerIP string field (old management). diff --git a/client/internal/acl/manager_test.go b/client/internal/acl/manager_test.go index 968654ae9..8f737706e 100644 --- a/client/internal/acl/manager_test.go +++ b/client/internal/acl/manager_test.go @@ -5,9 +5,9 @@ import ( "net/netip" "testing" - "github.com/golang/mock/gomock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" "github.com/netbirdio/netbird/client/firewall" "github.com/netbirdio/netbird/client/iface" @@ -87,7 +87,7 @@ func TestDefaultManager(t *testing.T) { networkMap.FirewallRules = append( networkMap.FirewallRules, &mgmProto.FirewallRule{ - PeerIP: "10.93.0.3", + PeerIP: "10.93.0.3", //nolint:staticcheck Direction: mgmProto.RuleDirection_IN, Action: mgmProto.RuleAction_DROP, Protocol: mgmProto.RuleProtocol_ICMP, @@ -556,12 +556,12 @@ func TestApplyFilteringSkipsUnchangedConfig(t *testing.T) { func buildNetworkMap(peerRules, routeRules int) *mgmProto.NetworkMap { nm := &mgmProto.NetworkMap{ - FirewallRulesIsEmpty: peerRules == 0, + FirewallRulesIsEmpty: peerRules == 0, RoutesFirewallRulesIsEmpty: routeRules == 0, } for i := range peerRules { nm.FirewallRules = append(nm.FirewallRules, &mgmProto.FirewallRule{ - PeerIP: fmt.Sprintf("10.%d.%d.%d", i>>16&0xff, i>>8&0xff, i&0xff), + PeerIP: fmt.Sprintf("10.%d.%d.%d", i>>16&0xff, i>>8&0xff, i&0xff), //nolint:staticcheck Direction: mgmProto.RuleDirection_IN, Action: mgmProto.RuleAction_ACCEPT, Protocol: mgmProto.RuleProtocol_TCP, diff --git a/client/internal/acl/mocks/iface_mapper.go b/client/internal/acl/mocks/iface_mapper.go index 95d5a2c58..f8cca1c2d 100644 --- a/client/internal/acl/mocks/iface_mapper.go +++ b/client/internal/acl/mocks/iface_mapper.go @@ -7,7 +7,7 @@ package mocks import ( reflect "reflect" - gomock "github.com/golang/mock/gomock" + gomock "go.uber.org/mock/gomock" wgdevice "golang.zx2c4.com/wireguard/device" "github.com/netbirdio/netbird/client/iface/device" diff --git a/client/internal/auth/auth.go b/client/internal/auth/auth.go index 51f56b644..b3a9e1158 100644 --- a/client/internal/auth/auth.go +++ b/client/internal/auth/auth.go @@ -138,26 +138,37 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) { // GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection // This avoids creating a new connection to the management server -func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool) (OAuthFlow, error) { +func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint string) (OAuthFlow, error) { var flow OAuthFlow - var err error - err = a.withRetry(ctx, func(client *mgm.GrpcClient) error { + err := a.withRetry(ctx, func(client *mgm.GrpcClient) error { if forceDeviceAuth { - flow, err = a.getDeviceFlow(client) - return err + deviceFlow, err := a.getDeviceFlow(client) + if err != nil { + return err + } + deviceFlow.SetLoginHint(hint) + flow = deviceFlow + return nil } // Try PKCE flow first - flow, err = a.getPKCEFlow(client) + pkceFlow, err := a.getPKCEFlow(client) if err != nil { // If PKCE not supported, try Device flow if s, ok := status.FromError(err); ok && (s.Code() == codes.NotFound || s.Code() == codes.Unimplemented) { - flow, err = a.getDeviceFlow(client) - return err + deviceFlow, err := a.getDeviceFlow(client) + if err != nil { + return err + } + deviceFlow.SetLoginHint(hint) + flow = deviceFlow + return nil } return err } + pkceFlow.SetLoginHint(hint) + flow = pkceFlow return nil }) @@ -351,6 +362,7 @@ func (a *Auth) setSystemInfoFlags(info *system.Info) { a.config.BlockLANAccess, a.config.BlockInbound, a.config.DisableIPv6, + a.config.SyncMessageVersion, a.config.EnableSSHRoot, a.config.EnableSSHSFTP, a.config.EnableSSHLocalPortForwarding, diff --git a/client/internal/auth/oauth.go b/client/internal/auth/oauth.go index a50a2ce6f..91329c98b 100644 --- a/client/internal/auth/oauth.go +++ b/client/internal/auth/oauth.go @@ -97,9 +97,7 @@ func authenticateWithPKCEFlow(ctx context.Context, config *profilemanager.Config return nil, fmt.Errorf("getting pkce authorization flow info failed with error: %v", err) } - if hint != "" { - pkceFlowInfo.SetLoginHint(hint) - } + pkceFlowInfo.SetLoginHint(hint) return pkceFlowInfo, nil } @@ -127,9 +125,7 @@ func authenticateWithDeviceCodeFlow(ctx context.Context, config *profilemanager. } } - if hint != "" { - deviceFlowInfo.SetLoginHint(hint) - } + deviceFlowInfo.SetLoginHint(hint) return deviceFlowInfo, nil } diff --git a/client/internal/auth/sessionwatch/watcher.go b/client/internal/auth/sessionwatch/watcher.go index e75a7022e..e685c28d0 100644 --- a/client/internal/auth/sessionwatch/watcher.go +++ b/client/internal/auth/sessionwatch/watcher.go @@ -24,11 +24,7 @@ import ( ) const ( - // Skew tolerates a small clock difference between the management - // server and this peer before treating a deadline as "in the past". - // Slightly above typical NTP drift; tight enough that the UI doesn't - // paint a stale expiry as if it were valid. - Skew = 30 * time.Second + maxPastHorizon = 30 * 24 * time.Hour // maxDeadlineHorizon caps how far in the future an accepted deadline // can sit. A timestamp beyond this is almost certainly a protocol @@ -57,7 +53,7 @@ var ( ErrDeadlineTooFarFuture = errors.New("session deadline too far in the future") // ErrDeadlineInPast is returned by Update when the supplied deadline - // is more than Skew in the past. + // is more than maxPastHorizon in the past. ErrDeadlineInPast = errors.New("session deadline in the past") ) @@ -66,15 +62,14 @@ var ( // for deadline change/clear, PublishEvent for the two warnings); tests pass // a fake recorder so the same surface is observable without an engine. // -// The watcher is the single owner of the deadline propagated to the -// recorder: every set, clear, sanity-check rejection and Close routes the -// value through SetSessionExpiresAt, so the SubscribeStatus snapshot the UI -// reads can never drift from the watcher's timer state. (SetSessionExpiresAt -// fans out its own state-change notification, so no separate notify is -// needed.) The recorder is server-scoped and outlives this engine-scoped -// watcher — without the Close-time clear a teardown (Down, or the Down+Up of -// a profile switch) would leave the next session showing the previous one's -// stale "expires in" value. +// While the watcher runs, it owns the deadline propagated to the recorder: +// every set, clear and sanity-check rejection routes the value through +// SetSessionExpiresAt, so the SubscribeStatus snapshot the UI reads can +// never drift from the watcher's timer state. (SetSessionExpiresAt fans +// out its own state-change notification, so no separate notify is needed.) +// The recorder is server-scoped and outlives this engine-scoped watcher; +// Close deliberately leaves the recorder value in place so transient engine +// restarts don't blank it — the client run loop clears it on real teardown. // // PublishEvent's signature mirrors peer.Status.PublishEvent: the watcher // composes the metadata internally so the wire format (MetaSession*) is @@ -135,10 +130,13 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher { // was disabled). // // Same-value updates are no-ops. A different non-zero value cancels any -// pending timer, resets the "already fired" guard, and arms a new one. +// pending timer, resets the "already fired" guards, and — when the +// deadline lies in the future — arms fresh warning timers. A deadline +// already in the past (within maxPastHorizon) is recorded as-is with no +// timers: the session has expired and consumers render it that way. // // Returns one of the sentinel Err* values when the deadline fails the -// sanity checks (pre-epoch, far future, or in the past beyond Skew). +// sanity checks (pre-epoch, far future, or past beyond maxPastHorizon). // In every error case the watcher first clears its state so it stays // consistent with what the caller will push into its other sinks (e.g. // applySessionDeadline forces a zero deadline into the status recorder @@ -163,7 +161,7 @@ func (w *Watcher) Update(deadline time.Time) error { case deadline.After(now.Add(maxDeadlineHorizon)): w.clearLocked() return fmt.Errorf("%w: %v", ErrDeadlineTooFarFuture, deadline) - case deadline.Before(now.Add(-Skew)): + case deadline.Before(now.Add(-maxPastHorizon)): w.clearLocked() return fmt.Errorf("%w: %v (now=%v)", ErrDeadlineInPast, deadline, now) } @@ -183,7 +181,9 @@ func (w *Watcher) Update(deadline time.Time) error { w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} - w.armTimerLocked(deadline) + if deadline.After(now) { + w.armTimerLocked(deadline) + } recorder := w.recorder w.mu.Unlock() if recorder != nil { @@ -227,30 +227,25 @@ func (w *Watcher) Dismiss() { log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339)) } -// Close stops any pending timer and drops the deadline on the status -// recorder. Update calls after Close are ignored. Clearing the recorder -// here is what keeps a teardown (Down, or the Down+Up of a profile switch) -// from leaving the next session showing this one's stale "expires in" -// value — the recorder is server-scoped and outlives this engine-scoped -// watcher, so nothing else drops the anchor on teardown. +// Close stops any pending timer. Update calls after Close are ignored. +// The recorder keeps its deadline: the watcher is engine-scoped and closes +// on every engine restart (network change, sleep/wake, stream errors) +// while the SSO deadline stays valid across those, so clearing here would +// blank the UI's "expires in" row on every transient reconnect. The +// client run loop clears the server-scoped recorder when it exits for +// real (Down, profile switch, permanent login failure). func (w *Watcher) Close() { w.mu.Lock() + defer w.mu.Unlock() if w.closed { - w.mu.Unlock() return } w.closed = true w.stopTimerLocked() - hadDeadline := !w.current.IsZero() w.current = time.Time{} w.firedAt = time.Time{} w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} - recorder := w.recorder - w.mu.Unlock() - if recorder != nil && hadDeadline { - recorder.SetSessionExpiresAt(time.Time{}) - } } // clearLocked drops the tracked deadline and notifies the recorder so diff --git a/client/internal/auth/sessionwatch/watcher_test.go b/client/internal/auth/sessionwatch/watcher_test.go index da2b6add6..4b49a94b6 100644 --- a/client/internal/auth/sessionwatch/watcher_test.go +++ b/client/internal/auth/sessionwatch/watcher_test.go @@ -224,11 +224,13 @@ func TestNewDeadlineCancelsPriorTimer(t *testing.T) { func TestRefreshAfterFireArmsNewWarning(t *testing.T) { r := &fakeRecorder{} - lead := 30 * time.Millisecond + lead := 150 * time.Millisecond w := newWatcher(lead, r) defer w.Close() - first := time.Now().Add(50 * time.Millisecond) + // Warning fires ~20ms in; the deadline itself stays 150ms away so the + // replacement below lands well before it. + first := time.Now().Add(170 * time.Millisecond) _ = w.Update(first) // Wait for stateChange + warning of the first cycle. @@ -306,7 +308,29 @@ func TestUpdateRejectsTooFarFuture(t *testing.T) { } } -func TestUpdateInPastClearsDeadline(t *testing.T) { +func TestUpdateRecentPastRecordedAsExpired(t *testing.T) { + r := &fakeRecorder{} + w := newWatcher(50*time.Millisecond, r) + defer w.Close() + + d := time.Now().Add(-1 * time.Hour) + if err := w.Update(d); err != nil { + t.Fatalf("recent-past Update should succeed, got %v", err) + } + if !w.Deadline().Equal(d) { + t.Fatalf("expected deadline to be recorded, got %v want %v", w.Deadline(), d) + } + if got := r.deadline(); !got.Equal(d) { + t.Fatalf("recorder deadline = %v, want %v", got, d) + } + + time.Sleep(80 * time.Millisecond) + if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 { + t.Fatalf("no warning events may fire for an already-past deadline, got %+v", r.snapshot()) + } +} + +func TestUpdateAncientPastRejected(t *testing.T) { r := &fakeRecorder{} w := newWatcher(50*time.Millisecond, r) defer w.Close() @@ -318,12 +342,12 @@ func TestUpdateInPastClearsDeadline(t *testing.T) { // Drain the stateChange from the seed. waitForEvents(t, r, 1) - err := w.Update(time.Now().Add(-1 * time.Hour)) + err := w.Update(time.Now().Add(-31 * 24 * time.Hour)) if !errors.Is(err, ErrDeadlineInPast) { t.Fatalf("want ErrDeadlineInPast, got %v", err) } if !w.Deadline().IsZero() { - t.Fatalf("in-past update must clear the deadline, got %v", w.Deadline()) + t.Fatalf("rejected ancient-past update must clear the deadline, got %v", w.Deadline()) } events := waitForEvents(t, r, 2) if events[1].kind != stateChange { @@ -331,39 +355,25 @@ func TestUpdateInPastClearsDeadline(t *testing.T) { } } -func TestUpdateWithinSkewAccepted(t *testing.T) { - r := &fakeRecorder{} - w := newWatcher(50*time.Millisecond, r) - defer w.Close() - - // 5 seconds in the past is within the 30s Skew tolerance — accept it. - d := time.Now().Add(-5 * time.Second) - if err := w.Update(d); err != nil { - t.Fatalf("within-skew Update should succeed, got %v", err) - } - if !w.Deadline().Equal(d) { - t.Fatalf("expected deadline to be applied, got %v want %v", w.Deadline(), d) - } -} - func TestCloseSilencesUpdates(t *testing.T) { r := &fakeRecorder{} w := newWatcher(50*time.Millisecond, r) w.Close() - _ = w.Update(time.Now().Add(time.Hour)) - - time.Sleep(20 * time.Millisecond) + if err := w.Update(time.Now().Add(time.Hour)); err != nil { + t.Fatalf("Update after Close: want nil, got %v", err) + } if got := r.snapshot(); len(got) != 0 { t.Fatalf("expected no events after Close, got %+v", got) } } -// TestCloseClearsRecorderDeadline pins the profile-switch fix: a watcher -// holding a live deadline must zero the recorder on Close so the next -// engine's watcher (and the UI reading the shared server-scoped recorder) -// doesn't start out showing the previous session's stale "expires in". -func TestCloseClearsRecorderDeadline(t *testing.T) { +// TestCloseKeepsRecorderDeadline pins the reconnect-flap fix: the watcher +// closes on every engine restart (network change, sleep/wake) while the +// SSO deadline stays valid across those, so Close must leave the +// server-scoped recorder's value in place. The client run loop clears the +// recorder when it exits for real. +func TestCloseKeepsRecorderDeadline(t *testing.T) { r := &fakeRecorder{} w := newWatcher(time.Hour, r) @@ -377,8 +387,8 @@ func TestCloseClearsRecorderDeadline(t *testing.T) { w.Close() - if got := r.deadline(); !got.IsZero() { - t.Fatalf("recorder deadline after Close = %v, want zero", got) + if got := r.deadline(); !got.Equal(d) { + t.Fatalf("recorder deadline after Close = %v, want %v", got, d) } } diff --git a/client/internal/conn_mgr.go b/client/internal/conn_mgr.go index 77d1e6ca5..ad0f00c5d 100644 --- a/client/internal/conn_mgr.go +++ b/client/internal/conn_mgr.go @@ -34,6 +34,8 @@ const ( // - Handling connection establishment based on peer signaling // // The implementation is not thread-safe; it is protected by engine.syncMsgMux. +// The only exception is ActivatePeer, which is safe for concurrent use so the +// DNS warm-up path can call it without contending on the engine mutex. type ConnMgr struct { peerStore *peerstore.Store statusRecorder *peer.Status @@ -42,12 +44,26 @@ type ConnMgr struct { rosenpassEnabled bool lazyConnMgr *manager.Manager + // lazyConnMgrMu guards the lazyConnMgr pointer for readers outside the + // engine loop (ActivatePeer). Writers hold it in addition to + // engine.syncMsgMux; all other reads stay under engine.syncMsgMux only. + lazyConnMgrMu sync.RWMutex + + // reconcileRoutedIPs re-applies a peer's routed allowed IPs after its lazy wake endpoint is + // (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile. + reconcileRoutedIPs func(peerKey string) error wg sync.WaitGroup lazyCtx context.Context lazyCtxCancel context.CancelFunc } +// SetRoutedIPsReconciler injects the callback used to re-apply a peer's routed allowed IPs when +// its lazy wake endpoint is (re)armed. Must be called before the lazy manager starts. +func (e *ConnMgr) SetRoutedIPsReconciler(fn func(peerKey string) error) { + e.reconcileRoutedIPs = fn +} + func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr { e := &ConnMgr{ peerStore: peerStore, @@ -238,12 +254,20 @@ func (e *ConnMgr) RemovePeerConn(peerKey string) { conn.Log.Infof("removed peer from lazy conn manager") } +// ActivatePeer wakes an idle lazy connection. Unlike the rest of ConnMgr it is +// safe for concurrent use: the lazy manager pointer is read under lazyConnMgrMu +// and the manager itself is internally synchronized, so callers outside the +// engine loop (DNS warm-up) do not need engine.syncMsgMux. func (e *ConnMgr) ActivatePeer(ctx context.Context, conn *peer.Conn) { - if !e.isStartedWithLazyMgr() { + e.lazyConnMgrMu.RLock() + lazyConnMgr := e.lazyConnMgr + started := lazyConnMgr != nil && e.lazyCtxCancel != nil + e.lazyConnMgrMu.RUnlock() + if !started { return } - if found := e.lazyConnMgr.ActivatePeer(conn.GetKey()); found { + if found := lazyConnMgr.ActivatePeer(conn.GetKey()); found { if err := conn.Open(ctx); err != nil { conn.Log.Errorf("failed to open connection: %v", err) } @@ -268,16 +292,22 @@ func (e *ConnMgr) Close() { e.lazyCtxCancel() e.wg.Wait() + + e.lazyConnMgrMu.Lock() e.lazyConnMgr = nil + e.lazyConnMgrMu.Unlock() } func (e *ConnMgr) initLazyManager(engineCtx context.Context) { cfg := manager.Config{ InactivityThreshold: inactivityThresholdEnv(), + ReconcileAllowedIPs: e.reconcileRoutedIPs, } - e.lazyConnMgr = manager.NewManager(cfg, engineCtx, e.peerStore, e.iface) + e.lazyConnMgrMu.Lock() + e.lazyConnMgr = manager.NewManager(cfg, engineCtx, e.peerStore, e.iface) e.lazyCtx, e.lazyCtxCancel = context.WithCancel(engineCtx) + e.lazyConnMgrMu.Unlock() e.wg.Add(1) go func() { @@ -316,7 +346,10 @@ func (e *ConnMgr) closeManager(ctx context.Context) { e.lazyCtxCancel() e.wg.Wait() + + e.lazyConnMgrMu.Lock() e.lazyConnMgr = nil + e.lazyConnMgrMu.Unlock() for _, peerID := range e.peerStore.PeersPubKey() { e.peerStore.PeerConnOpen(ctx, peerID) @@ -352,11 +385,20 @@ func inactivityThresholdEnv() *time.Duration { return nil } - parsedMinutes, err := strconv.Atoi(envValue) - if err != nil || parsedMinutes <= 0 { - return nil + // Documented format: a Go duration such as "30m" or "1h". + if d, err := time.ParseDuration(envValue); err == nil { + if d <= 0 { + return nil + } + return &d } - d := time.Duration(parsedMinutes) * time.Minute - return &d + // Backwards compatibility: a bare integer used to be interpreted as minutes. + if parsedMinutes, err := strconv.Atoi(envValue); err == nil && parsedMinutes > 0 { + d := time.Duration(parsedMinutes) * time.Minute + return &d + } + + log.Warnf("invalid %s value %q: expected a Go duration such as 30m or 1h", lazyconn.EnvInactivityThreshold, envValue) + return nil } diff --git a/client/internal/conn_mgr_test.go b/client/internal/conn_mgr_test.go index 5e2c53e35..ac5d6f2c8 100644 --- a/client/internal/conn_mgr_test.go +++ b/client/internal/conn_mgr_test.go @@ -1,10 +1,21 @@ package internal import ( + "context" + "net" + "net/netip" "os" + "sync" "testing" + "time" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" + + "github.com/netbirdio/netbird/client/iface/wgaddr" "github.com/netbirdio/netbird/client/internal/lazyconn" + "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/peerstore" + "github.com/netbirdio/netbird/monotime" ) func TestResolveLazyForce(t *testing.T) { @@ -38,3 +49,93 @@ func TestResolveLazyForce(t *testing.T) { }) } } + +type mockLazyWGIface struct{} + +func (mockLazyWGIface) RemovePeer(string) error { return nil } +func (mockLazyWGIface) UpdatePeer(string, []netip.Prefix, time.Duration, *net.UDPAddr, *wgtypes.Key) error { + return nil +} +func (mockLazyWGIface) IsUserspaceBind() bool { return false } +func (mockLazyWGIface) Address() wgaddr.Address { return wgaddr.Address{} } +func (mockLazyWGIface) LastActivities() map[string]monotime.Time { return nil } +func (mockLazyWGIface) MTU() uint16 { return 1280 } + +// TestConnMgr_ActivatePeerConcurrentWithLifecycle exercises ActivatePeer from +// non-engine goroutines (the DNS warm-up path) racing the manager lifecycle, +// which stays on the engine loop. Run with -race: it fails if ActivatePeer +// still requires engine.syncMsgMux for safety. +func TestConnMgr_ActivatePeerConcurrentWithLifecycle(t *testing.T) { + t.Setenv(lazyconn.EnvLazyConn, "on") + + status := peer.NewRecorder("https://mgm") + store := peerstore.NewConnStore() + connMgr := NewConnMgr(&EngineConfig{}, status, store, mockLazyWGIface{}) + + conn := newTestPeerConn(t, "peerA") + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + connMgr.Start(ctx) + + done := make(chan struct{}) + var wg sync.WaitGroup + for range 4 { + wg.Add(1) + go func() { + defer wg.Done() + for { + select { + case <-done: + return + default: + connMgr.ActivatePeer(ctx, conn) + } + } + }() + } + + // Let the activators spin against the started manager, then tear it down + // underneath them and let them spin against the stopped manager. + time.Sleep(100 * time.Millisecond) + connMgr.Close() + time.Sleep(50 * time.Millisecond) + + close(done) + wg.Wait() +} + +func TestInactivityThresholdEnv(t *testing.T) { + tests := []struct { + name string + val string + want *time.Duration + }{ + {name: "unset", val: "", want: nil}, + {name: "go duration minutes", val: "30m", want: durPtr(30 * time.Minute)}, + {name: "go duration hours", val: "1h", want: durPtr(time.Hour)}, + {name: "go duration seconds", val: "90s", want: durPtr(90 * time.Second)}, + {name: "bare integer is minutes (backwards compat)", val: "5", want: durPtr(5 * time.Minute)}, + {name: "zero duration", val: "0s", want: nil}, + {name: "zero integer", val: "0", want: nil}, + {name: "negative duration", val: "-5m", want: nil}, + {name: "garbage", val: "abc", want: nil}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Setenv(lazyconn.EnvInactivityThreshold, tc.val) + got := inactivityThresholdEnv() + switch { + case tc.want == nil && got != nil: + t.Fatalf("want nil, got %v", *got) + case tc.want != nil && got == nil: + t.Fatalf("want %v, got nil", *tc.want) + case tc.want != nil && *got != *tc.want: + t.Fatalf("want %v, got %v", *tc.want, *got) + } + }) + } +} + +func durPtr(d time.Duration) *time.Duration { return &d } diff --git a/client/internal/connect.go b/client/internal/connect.go index c2fc2fd73..ca50f912f 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -34,9 +34,11 @@ import ( "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/statemanager" "github.com/netbirdio/netbird/client/internal/stdnet" + "github.com/netbirdio/netbird/client/internal/tunnelnotifier" "github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/internal/updater/installer" nbnet "github.com/netbirdio/netbird/client/net" + "github.com/netbirdio/netbird/client/netevents" cProto "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/ssh" sshconfig "github.com/netbirdio/netbird/client/ssh/config" @@ -69,18 +71,31 @@ type ConnectClient struct { updateManager *updater.Manager persistSyncResponse bool + + // netMgr gates every reconnection loop on OS-reported network + // availability and sweeps connections on network change. + netMgr *netevents.Manager +} + +// ConnectClientOption configures optional ConnectClient behavior. +type ConnectClientOption func(*ConnectClient) + +// WithNetEvents injects the OS network event handling. +func WithNetEvents(events *netevents.Manager) ConnectClientOption { + return func(c *ConnectClient) { c.netMgr = events } } func NewConnectClient( ctx context.Context, config *profilemanager.Config, statusRecorder *peer.Status, + opts ...ConnectClientOption, ) *ConnectClient { // Derive the run context here so Stop owns the cancel that unblocks the run // loop. runCancel is set once at construction, so Stop can call it without // racing the run loop's startup. Callers therefore need not cancel before Stop. runCtx, runCancel := context.WithCancel(ctx) - return &ConnectClient{ + c := &ConnectClient{ ctx: runCtx, runCancel: runCancel, runExited: make(chan struct{}), @@ -88,6 +103,10 @@ func NewConnectClient( statusRecorder: statusRecorder, engineMutex: sync.Mutex{}, } + for _, opt := range opts { + opt(c) + } + return c } func (c *ConnectClient) SetUpdateManager(um *updater.Manager) { @@ -112,11 +131,14 @@ func (c *ConnectClient) RunOnAndroid( stateFilePath string, cacheDir string, ) error { + notifier := tunnelnotifier.New(networkChangeListener, nil) + defer notifier.Close() + // in case of non Android os these variables will be nil mobileDependency := MobileDependency{ TunAdapter: tunAdapter, IFaceDiscover: iFaceDiscover, - NetworkChangeListener: networkChangeListener, + NetworkChangeListener: notifier, HostDNSAddresses: dnsAddresses, DnsReadyListener: dnsReadyListener, StateFilePath: stateFilePath, @@ -136,10 +158,13 @@ func (c *ConnectClient) RunOniOS( // Set GC percent to 5% to reduce memory usage as iOS only allows 50MB of memory for the extension. debug.SetGCPercent(5) + notifier := tunnelnotifier.New(networkChangeListener, dnsManager) + defer notifier.Close() + mobileDependency := MobileDependency{ FileDescriptor: fileDescriptor, - NetworkChangeListener: networkChangeListener, - DnsManager: dnsManager, + NetworkChangeListener: notifier, + DnsManager: notifier, StateFilePath: stateFilePath, TempDir: cacheDir, } @@ -257,13 +282,23 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan log.Errorf("failed to clean up temporary installer file: %v", err) } - defer c.statusRecorder.ClientStop() + defer func() { + c.statusRecorder.SetSessionExpiresAt(time.Time{}) + c.statusRecorder.ClientStop() + }() operation := func() error { // if context cancelled we not start new backoff cycle if c.ctx.Err() != nil { return nil } + // suspend connection attempts while the OS reports no usable network + if waited, err := c.netMgr.Wait(c.ctx); err != nil { + return nil + } else if waited { + backOff.Reset() + } + state.Set(StatusConnecting) engineCtx, cancel := context.WithCancel(c.ctx) @@ -275,7 +310,8 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan }() log.Debugf("connecting to the Management service %s", c.config.ManagementURL.Host) - mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled) + mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled, + mgm.WithNetEvents(c.netMgr)) if err != nil { // On daemon shutdown / Down() the parent context is cancelled // and the dial fails with "context canceled". Wrapping that @@ -350,7 +386,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan }() // with the global Netbird config in hand connect (just a connection, no stream yet) Signal - signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey) + signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey, c.netMgr) if err != nil { log.Error(err) return wrapErr(err) @@ -386,7 +422,8 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan engineConfig.StateDir = filepath.Dir(path) } - relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU) + relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU, + relayClient.WithNetEvents(c.netMgr)) c.statusRecorder.SetRelayMgr(relayManager) if len(relayURLs) > 0 { if token != nil { @@ -414,6 +451,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan UpdateManager: c.updateManager, ClientMetrics: c.clientMetrics, MetricsCtx: c.ctx, + NetMgr: c.netMgr, }, mobileDependency) engine.SetSyncResponsePersistence(c.persistSyncResponse) c.engine = engine @@ -470,6 +508,16 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan // status stream stuck at Connecting. err = backoff.Retry(operation, backoff.WithContext(backOff, c.ctx)) if err != nil { + // Once the client context is cancelled backoff.WithContext surfaces the + // bare context error, and any attempt torn down mid-flight reports the + // same. That cancellation is the caller asking us to stop (Stop, Down or + // an engine restart), so exit cleanly instead of handing back a failure + // the caller would have to distinguish from a real one. + if c.ctx.Err() != nil && errors.Is(err, context.Canceled) { + log.Info("exiting client retry loop, context cancelled") + return nil + } + log.Debugf("exiting client retry loop due to unrecoverable error: %s", err) if s, ok := gstatus.FromError(err); ok && (s.Code() == codes.PermissionDenied) { state.Set(StatusNeedsLogin) @@ -618,6 +666,7 @@ func createEngineConfig(key wgtypes.Key, config *profilemanager.Config, peerConf BlockLANAccess: config.BlockLANAccess, BlockInbound: config.BlockInbound, DisableIPv6: config.DisableIPv6, + SyncMessageVersion: config.SyncMessageVersion, LazyConnection: lazyconn.ParseState(config.LazyConnection), @@ -662,7 +711,7 @@ func selectMTU(localMTU uint16, peerMTU int32) uint16 { } // connectToSignal creates Signal Service client and established a connection -func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key) (*signal.GrpcClient, error) { +func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key, netMgr *netevents.Manager) (*signal.GrpcClient, error) { var sigTLSEnabled bool if wtConfig.Signal.Protocol == mgmProto.HostConfig_HTTPS { sigTLSEnabled = true @@ -670,7 +719,8 @@ func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourP sigTLSEnabled = false } - signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled) + signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled, + signal.WithNetEvents(netMgr)) if err != nil { log.Errorf("error while connecting to the Signal Exchange Service %s: %s", wtConfig.Signal.Uri, err) return nil, gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Signal Service : %s", err) @@ -693,6 +743,7 @@ func loginToManagement(ctx context.Context, client mgm.Client, pubSSHKey []byte, config.BlockLANAccess, config.BlockInbound, config.DisableIPv6, + config.SyncMessageVersion, config.EnableSSHRoot, config.EnableSSHSFTP, config.EnableSSHLocalPortForwarding, diff --git a/client/internal/daemonaddr/owner.go b/client/internal/daemonaddr/owner.go new file mode 100644 index 000000000..c476f9ae6 --- /dev/null +++ b/client/internal/daemonaddr/owner.go @@ -0,0 +1,15 @@ +package daemonaddr + +// DaemonRunsAsSelf reports whether the daemon listening at addr runs as this very +// user. That is what makes an unprivileged daemon authorize this process for the +// changes it otherwise restricts to root or an administrator, so a client can tell +// up front whether those controls are usable instead of letting a save fail. +// +// It is answered from the ownership of the socket or pipe the daemon created, so it +// costs no round trip and needs no cooperation from the daemon. Ownership that +// cannot be read is reported as false, including for a TCP address, so a caller +// reading this as "the daemon would allow it" fails closed. The daemon remains the +// only thing that authorizes anything: this only decides what a client offers. +func DaemonRunsAsSelf(addr string) bool { + return daemonRunsAsSelf(addr) +} diff --git a/client/internal/daemonaddr/owner_unix.go b/client/internal/daemonaddr/owner_unix.go new file mode 100644 index 000000000..493e6528d --- /dev/null +++ b/client/internal/daemonaddr/owner_unix.go @@ -0,0 +1,40 @@ +//go:build !windows + +package daemonaddr + +import ( + "os" + "strings" + "syscall" + + log "github.com/sirupsen/logrus" +) + +// daemonRunsAsSelf compares the owner of the daemon's Unix socket with this +// process's uid. Root is not treated specially here: a root caller is privileged +// on its own merits, and a root-owned socket says nothing about the caller. +func daemonRunsAsSelf(addr string) bool { + path, ok := strings.CutPrefix(addr, "unix://") + if !ok { + return false + } + + info, err := os.Stat(path) + if err != nil { + log.Debugf("stat daemon socket %s: %v", path, err) + return false + } + + // Only a socket says anything about a daemon. A directory or a leftover + // regular file at that path is not one, and reading it as "the daemon runs as + // us" would offer controls the daemon then refuses. + if info.Mode()&os.ModeSocket == 0 { + return false + } + + stat, ok := info.Sys().(*syscall.Stat_t) + if !ok { + return false + } + return stat.Uid == uint32(os.Getuid()) +} diff --git a/client/internal/daemonaddr/owner_unix_test.go b/client/internal/daemonaddr/owner_unix_test.go new file mode 100644 index 000000000..363c7d95d --- /dev/null +++ b/client/internal/daemonaddr/owner_unix_test.go @@ -0,0 +1,62 @@ +//go:build !windows + +package daemonaddr + +import ( + "net" + "os" + "path/filepath" + "testing" +) + +// A socket this user created means the daemon runs as this user, which is the +// rootless case where the daemon delegates its authority to its own identity. +func TestDaemonRunsAsSelf_OwnSocket(t *testing.T) { + path := filepath.Join(t.TempDir(), "netbird.sock") + ln, err := net.Listen("unix", path) + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { + if err := ln.Close(); err != nil { + t.Logf("close listener: %v", err) + } + }) + + if !DaemonRunsAsSelf("unix://" + path) { + t.Error("a socket owned by this user must count as the daemon running as us") + } +} + +// Everything that is not a readable socket of ours has to answer false, because +// the caller reads a true as "the daemon would authorize me". +func TestDaemonRunsAsSelf_FailsClosed(t *testing.T) { + dir := t.TempDir() + + // A socket owned by another user, which is what a root-run daemon looks like + // to an unprivileged client. Only assertable when we are not root ourselves. + rootOwned := "unix:///var/run/netbird.sock" + if _, err := os.Stat("/var/run/netbird.sock"); err == nil && os.Getuid() != 0 { + if DaemonRunsAsSelf(rootOwned) { + t.Error("a socket owned by another user must not count as ours") + } + } + + for name, addr := range map[string]string{ + "missing socket": "unix://" + filepath.Join(dir, "absent.sock"), + "tcp address": "tcp://127.0.0.1:41731", + "named pipe": "npipe://netbird", + "empty": "", + "no scheme": filepath.Join(dir, "absent.sock"), + "directory": "unix://" + dir, + "unknown scheme": "http://localhost:8080", + "scheme only": "unix://", + "relative socket": "unix://netbird.sock", + } { + t.Run(name, func(t *testing.T) { + if DaemonRunsAsSelf(addr) { + t.Errorf("%q must not count as a daemon running as us", addr) + } + }) + } +} diff --git a/client/internal/daemonaddr/owner_windows.go b/client/internal/daemonaddr/owner_windows.go new file mode 100644 index 000000000..1cd2bba15 --- /dev/null +++ b/client/internal/daemonaddr/owner_windows.go @@ -0,0 +1,42 @@ +//go:build windows + +package daemonaddr + +import ( + "context" + "strings" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// daemonRunsAsSelf reads the owner of the daemon's pipe. A daemon running as the +// service account owns its pipe as LocalSystem, and an elevated one as +// BUILTIN\Administrators, so only a daemon the user started themselves matches. +func daemonRunsAsSelf(addr string) bool { + name, ok := strings.CutPrefix(addr, pipeScheme) + if !ok { + return false + } + + for _, path := range PipePaths(name) { + // Bounded: this runs on the UI's path for deciding which controls to + // offer, so a pipe that does not answer promptly must not stall it. A + // timeout leaves the caller unprivileged, which only disables controls. + ctx, cancel := context.WithTimeout(context.Background(), probeTimeout) + conn, err := dialPipe(ctx, path) + cancel() + if err != nil { + continue + } + + owned := ipcauth.PipeOwnedBySelf(conn) + if cerr := conn.Close(); cerr != nil { + log.Debugf("close daemon pipe %s after ownership check: %v", path, cerr) + } + return owned + } + + return false +} diff --git a/client/internal/daemonaddr/pipe.go b/client/internal/daemonaddr/pipe.go new file mode 100644 index 000000000..51815ef5e --- /dev/null +++ b/client/internal/daemonaddr/pipe.go @@ -0,0 +1,103 @@ +package daemonaddr + +import ( + "context" + "net" + "runtime" + "strings" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" +) + +const ( + // WindowsPipeAddr is the default daemon address on Windows. A named pipe + // carries the connecting process's token, which loopback TCP does not, so + // it is the only Windows transport on which the daemon can tell who is + // calling it. + WindowsPipeAddr = "npipe://netbird" + + // legacyWindowsAddr is the loopback-TCP address the Windows daemon used + // before named-pipe support. + legacyWindowsAddr = "tcp://127.0.0.1:41731" + + pipeScheme = "npipe://" + + // protectedPrefix is the NPFS namespace in which only LocalSystem and + // members of BUILTIN\Administrators may create a pipe. A daemon running as + // the service account creates its pipe there so that an unprivileged process + // cannot pre-create the name, which would keep the daemon from starting and + // leave callers talking to the squatter. Opening such a pipe needs no + // privilege, so unprivileged clients still reach the daemon. + protectedPrefix = `ProtectedPrefix\Administrators\` +) + +// DialTarget returns the gRPC dial target and transport options for a daemon +// 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())} + + if name, ok := strings.CutPrefix(addr, pipeScheme); ok { + paths := PipePaths(name) + opts = append(opts, grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { + return dialPipePaths(ctx, paths) + })) + return "passthrough:///netbird-daemon-pipe", opts + } + + return strings.TrimPrefix(addr, "tcp://"), opts +} + +// PipePath maps an npipe address name ("netbird", from "npipe://netbird") to a +// Windows named-pipe path (\\.\pipe\netbird). A fully qualified path is left as +// is. +func PipePath(name string) string { + if strings.HasPrefix(name, `\\`) { + return name + } + return `\\.\pipe\` + name +} + +// PipePaths returns the paths a daemon control pipe may live at for an npipe +// address name, in the order both sides must try them: the protected name first, +// then the plain one. +// +// The daemon serves the first it can create, which is the protected name when it +// runs as the service account and the plain one when it runs as an ordinary user, +// as it does in netstack mode. Clients therefore have to try both, and because a +// client cannot tell from the name alone who created the pipe, the plain name is +// only usable once the server's identity has been checked: see +// verifyPipeServer. +// +// A fully qualified path is what the operator asked for and is used as is. +func PipePaths(name string) []string { + if strings.HasPrefix(name, `\\`) { + return []string{name} + } + return []string{PipePath(protectedPrefix + name), PipePath(name)} +} + +// IsProtectedPipePath reports whether a pipe path is in the namespace only an +// administrator or LocalSystem can create in, which is what lets a client trust +// such a pipe from its name alone. +func IsProtectedPipePath(path string) bool { + return strings.HasPrefix(path, `\\.\pipe\`+protectedPrefix) +} + +// MigrateLegacy upgrades the pre-named-pipe Windows daemon address to the named +// pipe, reporting whether it rewrote the address. Existing installs persist the +// daemon address, so without this an upgraded daemon would keep listening on +// loopback TCP, where callers carry no identity and privileged operations would +// have to be refused for everyone. Only the exact legacy default is rewritten: +// a deliberately chosen custom address is left alone. +func MigrateLegacy(addr string) (string, bool) { + return migrateLegacyForOS(runtime.GOOS, addr) +} + +func migrateLegacyForOS(goos, addr string) (string, bool) { + if goos == "windows" && addr == legacyWindowsAddr { + return WindowsPipeAddr, true + } + return addr, false +} diff --git a/client/internal/daemonaddr/pipe_other.go b/client/internal/daemonaddr/pipe_other.go new file mode 100644 index 000000000..04e8e7331 --- /dev/null +++ b/client/internal/daemonaddr/pipe_other.go @@ -0,0 +1,15 @@ +//go:build !windows + +package daemonaddr + +import ( + "context" + "fmt" + "net" +) + +// dialPipePaths is Windows-only: no other platform serves the daemon on a named +// pipe. +func dialPipePaths(context.Context, []string) (net.Conn, error) { + return nil, fmt.Errorf("named pipes are only supported on Windows") +} diff --git a/client/internal/daemonaddr/pipe_test.go b/client/internal/daemonaddr/pipe_test.go new file mode 100644 index 000000000..b9dfd90f1 --- /dev/null +++ b/client/internal/daemonaddr/pipe_test.go @@ -0,0 +1,30 @@ +package daemonaddr + +import ( + "slices" + "testing" +) + +// The protected name must be tried before the plain one on both sides: it is the +// one an unprivileged process cannot create, so preferring it is what keeps a +// squatter from owning the name the service daemon would otherwise use. +func TestPipePaths_PrefersTheProtectedName(t *testing.T) { + got := PipePaths("netbird") + want := []string{ + `\\.\pipe\ProtectedPrefix\Administrators\netbird`, + `\\.\pipe\netbird`, + } + if !slices.Equal(got, want) { + t.Errorf("PipePaths = %q, want %q", got, want) + } +} + +// An operator who passes a full path chose exactly one pipe, so neither side may +// look anywhere else. +func TestPipePaths_QualifiedPathIsUsedAsIs(t *testing.T) { + path := `\\.\pipe\custom-netbird` + got := PipePaths(path) + if !slices.Equal(got, []string{path}) { + t.Errorf("PipePaths = %q, want just %q", got, path) + } +} diff --git a/client/internal/daemonaddr/pipe_windows.go b/client/internal/daemonaddr/pipe_windows.go new file mode 100644 index 000000000..3cd10a6c3 --- /dev/null +++ b/client/internal/daemonaddr/pipe_windows.go @@ -0,0 +1,59 @@ +//go:build windows + +package daemonaddr + +import ( + "context" + "errors" + "fmt" + "net" + + "github.com/Microsoft/go-winio" + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// dialPipePaths connects to the first path that answers with a pipe server this +// client may trust, and returns the last error when none does. +func dialPipePaths(ctx context.Context, paths []string) (net.Conn, error) { + var lastErr error + for _, path := range paths { + conn, err := dialPipe(ctx, path) + if err != nil { + log.Debugf("dial daemon pipe %s: %v", path, err) + lastErr = err + continue + } + + // A pipe in the protected namespace could only have been created by an + // administrator or LocalSystem, so its name is the guarantee. Any other + // name has to be checked, because any local user can create one. + if !IsProtectedPipePath(path) { + if err := ipcauth.PipeServerTrusted(conn); err != nil { + if closeErr := conn.Close(); closeErr != nil { + log.Debugf("close untrusted pipe %s: %v", path, closeErr) + } + lastErr = fmt.Errorf("%s: %w", path, err) + continue + } + } + + return conn, nil + } + + if lastErr == nil { + lastErr = errors.New("no daemon pipe to connect to") + } + return nil, lastErr +} + +// dialPipe connects to the daemon control pipe at SECURITY_IDENTIFICATION. +// winio's plain DialPipe connects at SECURITY_ANONYMOUS, under which the daemon +// cannot read the caller's token at all. Identification lets the daemon read the +// caller's SID and groups without granting it the ability to act as the caller. +func dialPipe(ctx context.Context, path string) (net.Conn, error) { + access := uint32(windows.GENERIC_READ | windows.GENERIC_WRITE) + return winio.DialPipeAccessImpLevel(ctx, path, access, winio.PipeImpLevelIdentification) +} diff --git a/client/internal/daemonaddr/resolve_pipe_other.go b/client/internal/daemonaddr/resolve_pipe_other.go new file mode 100644 index 000000000..1aede8453 --- /dev/null +++ b/client/internal/daemonaddr/resolve_pipe_other.go @@ -0,0 +1,9 @@ +//go:build !windows + +package daemonaddr + +// ResolveDaemonAddr is a no-op off Windows, where there is no named-pipe +// default to fall back from. +func ResolveDaemonAddr(addr string) string { + return addr +} diff --git a/client/internal/daemonaddr/resolve_pipe_windows.go b/client/internal/daemonaddr/resolve_pipe_windows.go new file mode 100644 index 000000000..d12ddb15d --- /dev/null +++ b/client/internal/daemonaddr/resolve_pipe_windows.go @@ -0,0 +1,82 @@ +//go:build windows + +package daemonaddr + +import ( + "net" + "strings" + "time" + + "github.com/Microsoft/go-winio" + log "github.com/sirupsen/logrus" +) + +// probeTimeout bounds each transport probe. Both are local, so a daemon that is +// listening answers immediately and one that is not fails immediately. +const probeTimeout = 300 * time.Millisecond + +// ResolveDaemonAddr keeps a client on the named pipe and never silently moves it +// off. When the pipe does not answer it checks the legacy loopback TCP address, so +// a client meeting a daemon that has not restarted since the upgrade can say what +// is wrong, but it does not connect there. +// +// Using that address automatically would be a downgrade the user never asked for: +// any local process can bind 127.0.0.1 while the daemon is not listening, and the +// transport carries no caller identity, so a client that accepted whatever answered +// would hand a setup key, a pre-shared key or an SSO prompt to a local impostor. An +// operator who needs the legacy address during the upgrade window can still pass +// --daemon-addr explicitly, which is a deliberate choice and still refuses the +// privileged operations. +// +// Only the pipe address is resolved. A custom address is left alone, though passing +// --daemon-addr npipe://netbird explicitly is indistinguishable from the default +// here, so it is treated the same way. +func ResolveDaemonAddr(addr string) string { + if addr != WindowsPipeAddr { + return addr + } + + for _, path := range PipePaths("netbird") { + if pipeAvailable(path) { + return addr + } + } + + if tcpAvailable(legacyWindowsAddr) { + log.Warnf("the daemon is not serving %s, but something is listening on the legacy %s. "+ + "Restart the NetBird service so it serves the pipe. That address is not used automatically: "+ + "any local user can bind it and it carries no caller identity, so pass --daemon-addr %s "+ + "explicitly if you accept that", + WindowsPipeAddr, legacyWindowsAddr, legacyWindowsAddr) + } + + return addr +} + +func pipeAvailable(path string) bool { + timeout := probeTimeout + conn, err := winio.DialPipe(path, &timeout) + if err != nil { + return false + } + if err := conn.Close(); err != nil { + log.Debugf("close daemon pipe probe: %v", err) + } + return true +} + +func tcpAvailable(addr string) bool { + host := addr + if _, after, ok := strings.Cut(addr, "://"); ok { + host = after + } + + conn, err := net.DialTimeout("tcp", host, probeTimeout) + if err != nil { + return false + } + if err := conn.Close(); err != nil { + log.Debugf("close daemon TCP probe: %v", err) + } + return true +} diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go index 0e506ccd7..1d31c75ca 100644 --- a/client/internal/debug/debug.go +++ b/client/internal/debug/debug.go @@ -34,9 +34,8 @@ import ( "github.com/netbirdio/netbird/shared/netiputil" ) -const readmeContent = `Netbird debug bundle -This debug bundle contains the following files. -If the --anonymize flag is set, the files are anonymized to protect sensitive information. +const readmeContent = `This debug bundle contains the following files. +If anonymization is enabled (--anonymize / --anonymize-level), the files are anonymized to protect sensitive information. status.txt: Anonymized status information of the NetBird client. client.log: Most recent, anonymized client log file of the NetBird client. @@ -52,6 +51,7 @@ nftables.txt: Anonymized nftables rules with packet counters across all families sysctls.txt: Forwarding, reverse-path filter, source-validation, and conntrack accounting sysctl values that the NetBird client may read or modify, if --system-info flag was provided (Linux only). resolv.conf: DNS resolver configuration from /etc/resolv.conf (Unix systems only), if --system-info flag was provided. scutil_dns.txt: DNS configuration from scutil --dns (macOS only), if --system-info flag was provided. +dns_windows.txt: Anonymized NRPT rules and policy table in effect, DNS client policy, and per-interface and per-adapter DNS configuration (Windows only), if --system-info flag was provided. resolved_domains.txt: Anonymized resolved domain IP addresses from the status recorder. config.txt: Anonymized configuration information of the NetBird client. network_map.json: Anonymized sync response containing peer configurations, routes, DNS settings, and firewall rules. @@ -70,21 +70,34 @@ capture.pcap: Packet capture in pcap format. Only present when capture was runni Anonymization Process -The files in this bundle have been anonymized to protect sensitive information. Here's how the anonymization was applied: +The files in this bundle have been anonymized to protect sensitive information. The level applied to this bundle is recorded at the top of this file. Here's how the anonymization was applied: IP Addresses -IPv4 addresses are replaced with addresses starting from 198.51.100.0 -IPv6 addresses are replaced with addresses starting from 100:: +Default level: +- Public IPv4 addresses are replaced with addresses starting from 198.51.100.0 +- Public IPv6 addresses are replaced with addresses starting from 2001:db8:ffff:: +- IPv6 unique local addresses (fc00::/7) are anonymized as well: their random global ID uniquely identifies the network. +- IP addresses from internal IPv4 ranges and well-known addresses are not anonymized (e.g. 8.8.8.8, 100.64.0.0/10, addresses starting with 192.168., 172.16., 10., 169.254., fe80::). + +Strict level (--anonymize-level strict), in addition to the default level: +- Private (RFC 1918), CGNAT (100.64.0.0/10), and link-local (169.254.0.0/16, fe80::/10) addresses are anonymized too. +- Internal IPv4 addresses are replaced with addresses starting from 198.18.0.0 and internal IPv6 addresses with addresses starting from 2001:db8:1::, so internal addresses remain distinguishable from public ones. +- Addresses are mapped in order of first appearance: subnet structure, allocation scheme, and gateway conventions are not preserved. Prefix lengths of networks are preserved. +- Peer names in front of NetBird domains are replaced with numbered placeholders (e.g. peer-1.netbird.cloud), and subdomain labels of other domains with host-N placeholders. +- WireGuard public keys are replaced with consistent placeholder keys. -IP addresses from non public ranges and well known addresses are not anonymized (e.g. 8.8.8.8, 100.64.0.0/10, addresses starting with 192.168., 172.16., 10., etc.). Reoccuring IP addresses are replaced with the same anonymized address. Note: The anonymized IP addresses in the status file do not match those in the log and routes files. However, the anonymized IP addresses are consistent within the status file and across the routes and log files. +MAC Addresses +MAC addresses are replaced at every anonymization level with consistent placeholders counting up from 02:00:00:00:00:01. Broadcast, multicast, and all-zero addresses are kept. At the default level a preserved IPv6 link-local address may still embed a MAC address (EUI-64); the strict level anonymizes those addresses. + Domains All domain names (except for the netbird domains) are replaced with randomly generated strings ending in ".domain". Anonymized domains are consistent across all files in the bundle. Reoccuring domain names are replaced with the same anonymized domain. +At the strict level, the peer name labels in front of netbird domains are anonymized as well. Sync Response The network_map.json file contains the following anonymized information: @@ -225,11 +238,17 @@ scutil_dns.txt (macOS only): - Shows DNS configuration for all network interfaces - Includes search domains, nameservers, and DNS resolver settings - All IP addresses and domain names are anonymized + +dns_windows.txt (Windows only): +- Lists the NRPT rules of both policy stores, the local one and the group policy one, marking the rules the client created +- Follows them with the policy table the resolver has loaded, which differs from the rules while a change has not been picked up yet +- Includes the DNS client group policy, the global TCP/IP and Dnscache parameters, and the DNS values of every interface that has any +- Ends with the resolver configuration in effect per adapter, from GetAdaptersAddresses +- All IP addresses and domain names are anonymized ` const ( clientLogFile = "client.log" - uiLogFile = "gui-client.log" errorLogFile = "netbird.err" stdoutLogFile = "netbird.out" @@ -248,6 +267,20 @@ type MetricsExporter interface { Export(w io.Writer) error } +// LogOpener opens a log file for inclusion in the bundle. It exists so that log +// files whose path was supplied by an IPC caller can be opened under a check +// the daemon defines, instead of being opened with the daemon's privileges +// unconditionally. +type LogOpener func(path string) (*os.File, error) + +func openLogFile(path string) (*os.File, error) { + f, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + return f, nil +} + type BundleGenerator struct { anonymizer *anonymize.Anonymizer @@ -257,6 +290,7 @@ type BundleGenerator struct { syncResponse *mgmProto.SyncResponse logPath string uiLogPath string + uiLogOpener LogOpener tempDir string statePath string cpuProfile []byte @@ -267,6 +301,7 @@ type BundleGenerator struct { cliVersion string anonymize bool + anonymizeLevel anonymize.Level includeSystemInfo bool logFileCount uint32 @@ -274,7 +309,10 @@ type BundleGenerator struct { } type BundleConfig struct { - Anonymize bool + Anonymize bool + // AnonymizeLevel selects how much the anonymizer redacts. + // anonymize.LevelStrict implies Anonymize. + AnonymizeLevel anonymize.Level IncludeSystemInfo bool LogFileCount uint32 } @@ -285,14 +323,20 @@ type GeneratorDependencies struct { SyncResponse *mgmProto.SyncResponse LogPath string UILogPath string // Absolute path to the desktop UI's gui-client.log, reported via RegisterUILog. Empty if no UI registered one. - TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used. - StatePath string // Path to the state file. If empty, the ServiceManager default path is used. - CPUProfile []byte - CapturePath string - RefreshStatus func() - ClientMetrics MetricsExporter - DaemonVersion string - CliVersion string + // UILogOpener opens the UI log and its rotated siblings. The path comes from + // a local IPC caller, so the daemon must not open it with plain os.Open: the + // opener is where the caller's right to that file is enforced. Defaults to + // os.Open, which is only correct where the path is not caller-supplied + // (mobile). + UILogOpener LogOpener + TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used. + StatePath string // Path to the state file. If empty, the ServiceManager default path is used. + CPUProfile []byte + CapturePath string + RefreshStatus func() + ClientMetrics MetricsExporter + DaemonVersion string + CliVersion string } func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGenerator { @@ -302,14 +346,23 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen logFileCount = 1 } + uiLogOpener := deps.UILogOpener + if uiLogOpener == nil { + uiLogOpener = openLogFile + } + + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(cfg.AnonymizeLevel) + return &BundleGenerator{ - anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()), + anonymizer: anonymizer, internalConfig: deps.InternalConfig, statusRecorder: deps.StatusRecorder, syncResponse: deps.SyncResponse, logPath: deps.LogPath, uiLogPath: deps.UILogPath, + uiLogOpener: uiLogOpener, tempDir: deps.TempDir, statePath: deps.StatePath, cpuProfile: deps.CPUProfile, @@ -319,7 +372,8 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen daemonVersion: deps.DaemonVersion, cliVersion: deps.CliVersion, - anonymize: cfg.Anonymize, + anonymize: cfg.Anonymize || cfg.AnonymizeLevel >= anonymize.LevelStrict, + anonymizeLevel: cfg.AnonymizeLevel, includeSystemInfo: cfg.IncludeSystemInfo, logFileCount: logFileCount, } @@ -459,7 +513,13 @@ func (g *BundleGenerator) addSystemInfo() { } func (g *BundleGenerator) addReadme() error { - readmeReader := strings.NewReader(readmeContent) + level := "none (anonymization disabled)" + if g.anonymize { + level = g.anonymizeLevel.String() + } + header := fmt.Sprintf("Netbird debug bundle\nAnonymization level applied to this bundle: %s\n", level) + + readmeReader := strings.NewReader(header + readmeContent) if err := g.addFileToZip(readmeReader, "README.txt"); err != nil { return fmt.Errorf("add README file to zip: %w", err) } @@ -481,9 +541,10 @@ func (g *BundleGenerator) addStatus() error { fullStatus := g.statusRecorder.GetFullStatus() protoFullStatus := nbstatus.ToProtoFullStatus(fullStatus) overview := nbstatus.ConvertToStatusOutputOverview(protoFullStatus, nbstatus.ConvertOptions{ - Anonymize: g.anonymize, - ProfileName: profName, - DaemonVersion: g.daemonVersion, + Anonymize: g.anonymize, + AnonymizeLevel: g.anonymizeLevel, + ProfileName: profName, + DaemonVersion: g.daemonVersion, }) overview.CliVersion = g.cliVersion statusOutput := overview.FullDetailSummary() @@ -636,7 +697,7 @@ func (g *BundleGenerator) addCommonConfigFields(configContent *strings.Builder) configContent.WriteString("NetBird Client Configuration:\n\n") if key, err := wgtypes.ParseKey(g.internalConfig.PrivateKey); err == nil { - configContent.WriteString(fmt.Sprintf("PublicKey: %s\n", key.PublicKey().String())) + configContent.WriteString(fmt.Sprintf("PublicKey: %s\n", g.anonymizer.AnonymizeWGKey(key.PublicKey().String()))) } configContent.WriteString(fmt.Sprintf("WgIface: %s\n", g.internalConfig.WgIface)) configContent.WriteString(fmt.Sprintf("WgPort: %d\n", g.internalConfig.WgPort)) @@ -676,6 +737,7 @@ func (g *BundleGenerator) addCommonConfigFields(configContent *strings.Builder) configContent.WriteString(fmt.Sprintf("BlockLANAccess: %v\n", g.internalConfig.BlockLANAccess)) configContent.WriteString(fmt.Sprintf("BlockInbound: %v\n", g.internalConfig.BlockInbound)) configContent.WriteString(fmt.Sprintf("DisableIPv6: %v\n", g.internalConfig.DisableIPv6)) + configContent.WriteString(fmt.Sprintf("SyncMessageVersion: %v\n", g.internalConfig.SyncMessageVersion)) if g.internalConfig.DisableNotifications != nil { configContent.WriteString(fmt.Sprintf("DisableNotifications: %v\n", *g.internalConfig.DisableNotifications)) @@ -925,6 +987,11 @@ func (g *BundleGenerator) addUpdateLogs() error { } baseName := filepath.Base(logFile) + data, err = g.anonymizeBytes(data) + if err != nil { + log.Warnf("skipping update log file %s: %v", baseName, err) + continue + } if err := g.addFileToZip(bytes.NewReader(data), filepath.Join("update-logs", baseName)); err != nil { return fmt.Errorf("add update log file %s to zip: %w", baseName, err) } @@ -952,6 +1019,13 @@ func (g *BundleGenerator) addCorruptedStateFiles() error { } fileName := filepath.Base(match) + // Corrupted state files usually fail structured JSON anonymization, + // so run them through the string anonymizer instead. + data, err = g.anonymizeBytes(data) + if err != nil { + log.Warnf("skipping corrupted state file %s: %v", fileName, err) + continue + } if err := g.addFileToZip(bytes.NewReader(data), "corrupted_states/"+fileName); err != nil { log.Warnf("Failed to add corrupted state file %s to zip: %v", fileName, err) continue @@ -963,6 +1037,27 @@ func (g *BundleGenerator) addCorruptedStateFiles() error { return nil } +// anonymizeBytes runs raw file content through the string anonymizer line by +// line when anonymization is enabled. It errors instead of returning partial +// content, so a caller never adds an unanonymized fallback to the bundle. +func (g *BundleGenerator) anonymizeBytes(data []byte) ([]byte, error) { + if !g.anonymize { + return data, nil + } + + var buf bytes.Buffer + scanner := bufio.NewScanner(bytes.NewReader(data)) + scanner.Buffer(make([]byte, 1024*1024), 1024*1024) + for scanner.Scan() { + buf.WriteString(g.anonymizer.AnonymizeString(scanner.Text())) + buf.WriteByte('\n') + } + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("anonymize content: %w", err) + } + return buf.Bytes(), nil +} + func (g *BundleGenerator) addMetrics() error { if g.clientMetrics == nil { log.Debugf("skipping metrics in debug bundle: no metrics collector") @@ -995,11 +1090,11 @@ func (g *BundleGenerator) addLogfile() error { logDir := filepath.Dir(g.logPath) - if err := g.addSingleLogfile(g.logPath, clientLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, g.logPath, clientLogFile); err != nil { return fmt.Errorf("add client log file to zip: %w", err) } - g.addRotatedLogFiles(logDir, clientLogPrefix) + g.addRotatedLogFiles(openLogFile, logDir, clientLogPrefix) stdErrLogPath := filepath.Join(logDir, errorLogFile) stdoutLogPath := filepath.Join(logDir, stdoutLogFile) @@ -1008,11 +1103,11 @@ func (g *BundleGenerator) addLogfile() error { stdoutLogPath = darwinStdoutLogPath } - if err := g.addSingleLogfile(stdErrLogPath, errorLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, stdErrLogPath, errorLogFile); err != nil { log.Warnf("Failed to add %s to zip: %v", errorLogFile, err) } - if err := g.addSingleLogfile(stdoutLogPath, stdoutLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, stdoutLogPath, stdoutLogFile); err != nil { log.Warnf("Failed to add %s to zip: %v", stdoutLogFile, err) } @@ -1029,18 +1124,18 @@ func (g *BundleGenerator) addUILog() error { return nil } - if err := g.addSingleLogfile(g.uiLogPath, uiLogFile); err != nil { + if err := g.addSingleLogfile(g.uiLogOpener, g.uiLogPath, configs.UILogFile); err != nil { return fmt.Errorf("add UI log file to zip: %w", err) } - g.addRotatedLogFiles(filepath.Dir(g.uiLogPath), uiLogPrefix) + g.addRotatedLogFiles(g.uiLogOpener, filepath.Dir(g.uiLogPath), uiLogPrefix) return nil } // addSingleLogfile adds a single log file to the archive -func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error { - logFile, err := os.Open(logPath) +func (g *BundleGenerator) addSingleLogfile(open LogOpener, logPath, targetName string) error { + logFile, err := open(logPath) if err != nil { return fmt.Errorf("open log file %s: %w", targetName, err) } @@ -1065,8 +1160,8 @@ func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error { } // addSingleLogFileGz adds a single gzipped log file to the archive -func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error { - f, err := os.Open(logPath) +func (g *BundleGenerator) addSingleLogFileGz(open LogOpener, logPath, targetName string) error { + f, err := open(logPath) if err != nil { return fmt.Errorf("open gz log file %s: %w", targetName, err) } @@ -1113,7 +1208,7 @@ func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error { // addRotatedLogFiles adds rotated log files to the bundle based on logFileCount. // prefix is the base log name without extension (e.g. "client", "gui-client"); // the glob matches both files rotated by us and by logrotate on linux. -func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) { +func (g *BundleGenerator) addRotatedLogFiles(open LogOpener, logDir, prefix string) { if g.logFileCount == 0 { return } @@ -1153,9 +1248,9 @@ func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) { for i := 0; i < maxFiles; i++ { name := filepath.Base(files[i]) if strings.HasSuffix(name, ".gz") { - err = g.addSingleLogFileGz(files[i], name) + err = g.addSingleLogFileGz(open, files[i], name) } else { - err = g.addSingleLogfile(files[i], name) + err = g.addSingleLogfile(open, files[i], name) } if err != nil { log.Warnf("failed to add rotated log %s: %v", name, err) @@ -1435,6 +1530,7 @@ func anonymizeRemotePeer(peer *mgmProto.RemotePeerConfig, anonymizer *anonymize. } peer.Fqdn = anonymizer.AnonymizeDomain(peer.Fqdn) + peer.WgPubKey = anonymizer.AnonymizeWGKey(peer.WgPubKey) anonymizeSSHConfig(peer.SshConfig) } diff --git a/client/internal/debug/debug_ios.go b/client/internal/debug/debug_ios.go index a07c23dbd..001d64241 100644 --- a/client/internal/debug/debug_ios.go +++ b/client/internal/debug/debug_ios.go @@ -27,7 +27,7 @@ func (g *BundleGenerator) addPlatformLog() error { } swiftLogPath := filepath.Join(filepath.Dir(g.logPath), swiftLogFile) - if err := g.addSingleLogfile(swiftLogPath, swiftLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, swiftLogPath, swiftLogFile); err != nil { // The Swift log is best-effort: the app may not have written it yet. log.Warnf("failed to add %s to debug bundle: %v", swiftLogFile, err) } diff --git a/client/internal/debug/debug_linux.go b/client/internal/debug/debug_linux.go index 40d864eda..a36c0c0e7 100644 --- a/client/internal/debug/debug_linux.go +++ b/client/internal/debug/debug_linux.go @@ -844,6 +844,10 @@ func collectSysctls() string { []string{"net.ipv4.conf.all.src_valid_mark", "net.ipv4.conf.default.src_valid_mark"}, listInterfaceSysctls("ipv4", "src_valid_mark")..., )) + writeSysctlGroup(&builder, "accept_ra", append( + []string{"net.ipv6.conf.all.accept_ra", "net.ipv6.conf.default.accept_ra"}, + listInterfaceSysctls("ipv6", "accept_ra")..., + )) writeSysctlGroup(&builder, "conntrack", []string{ "net.netfilter.nf_conntrack_acct", "net.netfilter.nf_conntrack_tcp_loose", diff --git a/client/internal/debug/debug_logfiles_test.go b/client/internal/debug/debug_logfiles_test.go index 3749b3e9b..31420711f 100644 --- a/client/internal/debug/debug_logfiles_test.go +++ b/client/internal/debug/debug_logfiles_test.go @@ -97,7 +97,7 @@ func runAddRotatedLogFilesPrefix(t *testing.T, dir, prefix string, logFileCount archive: zip.NewWriter(&buf), logFileCount: logFileCount, } - g.addRotatedLogFiles(dir, prefix) + g.addRotatedLogFiles(openLogFile, dir, prefix) require.NoError(t, g.archive.Close()) zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len())) diff --git a/client/internal/debug/debug_nonunix.go b/client/internal/debug/debug_nonunix.go index 18d017050..adc9b9649 100644 --- a/client/internal/debug/debug_nonunix.go +++ b/client/internal/debug/debug_nonunix.go @@ -1,4 +1,4 @@ -//go:build !unix +//go:build !unix && !windows package debug diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go index 8286f6852..7fe93a5c1 100644 --- a/client/internal/debug/debug_test.go +++ b/client/internal/debug/debug_test.go @@ -887,6 +887,8 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) { ClientCertKeyPath: "/tmp/key", LazyConnection: "on", MTU: 1280, + DisableIPv6: true, + SyncMessageVersion: func(v int) *int { return &v }(1), } for _, anonymize := range []bool{false, true} { diff --git a/client/internal/debug/debug_windows.go b/client/internal/debug/debug_windows.go new file mode 100644 index 000000000..e88940fd3 --- /dev/null +++ b/client/internal/debug/debug_windows.go @@ -0,0 +1,443 @@ +//go:build windows + +package debug + +import ( + "encoding/hex" + "errors" + "fmt" + "net/netip" + "strings" + "unsafe" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" + "golang.org/x/sys/windows/registry" + + nbdns "github.com/netbirdio/netbird/client/internal/dns" +) + +const dnsInfoFileName = "dns_windows.txt" + +const ( + gpoDNSClientRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient` + tcpipParamsPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters` + dnscacheParams = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters` +) + +// interfaceDNSValues are the per-interface values that decide how a name is +// resolved and registered. Everything the DNS host manager writes is in here, +// so a bundle shows both what we set and what it replaced. +var interfaceDNSValues = []string{ + "NameServer", + "DhcpNameServer", + "Domain", + "DhcpDomain", + "SearchList", + "RegistrationEnabled", + "DisableDynamicUpdate", + "MaxNumberOfAddressesToRegister", + "EnableDHCP", +} + +// addDNSInfo collects and adds DNS configuration information to the archive +func (g *BundleGenerator) addDNSInfo() error { + if err := g.addFileToZip(strings.NewReader(g.collectDNSInfo()), dnsInfoFileName); err != nil { + return fmt.Errorf("add DNS info to zip: %w", err) + } + + return nil +} + +// collectDNSInfo renders the report. Everything below it reaches the platform +// through COM and through lazily resolved procedures, which panic when a +// procedure is missing rather than returning an error, and a debug bundle is not +// allowed to take the daemon down. The panic is contained here, and whatever was +// collected before it is kept and reported with it. +func (g *BundleGenerator) collectDNSInfo() (content string) { + var sb strings.Builder + + defer func() { + if r := recover(); r != nil { + log.Errorf("collecting Windows DNS configuration panicked: %v", r) + fmt.Fprintf(&sb, "\nerror: collection stopped: %v\n", r) + } + content = sb.String() + }() + + sb.WriteString("Windows DNS configuration\n") + sb.WriteString("=========================\n") + + adapters, adaptersErr := adapterAddresses() + + g.writeNRPTRules(&sb, "NRPT rules, local policy store", nbdns.DNSPolicyConfigRoot) + g.writeNRPTRules(&sb, "NRPT rules, group policy store", nbdns.GPODNSPolicyConfigRoot) + g.writeEffectiveNRPTPolicies(&sb) + g.writeRegistryKey(&sb, "DNS client group policy", gpoDNSClientRoot) + g.writeRegistryKey(&sb, "Global TCP/IP parameters", tcpipParamsPath) + g.writeRegistryKey(&sb, "Dnscache parameters", dnscacheParams) + g.writeInterfaceDNS(&sb, "Per-interface DNS, IPv4", nbdns.InterfaceConfigPath, adapterNames(adapters)) + g.writeInterfaceDNS(&sb, "Per-interface DNS, IPv6", nbdns.InterfaceConfigPathV6, adapterNames(adapters)) + g.writeAdapterDNS(&sb, adapters, adaptersErr) + + return sb.String() +} + +// writeNRPTRules lists every rule in a policy store, ours and any other +// product's, since a foreign rule for the same namespace decides resolution +// just as ours does. Rules the client wrote are marked. +func (g *BundleGenerator) writeNRPTRules(sb *strings.Builder, title, root string) { + writeSection(sb, title, root) + + names, err := subKeyNames(root) + if err != nil { + fmt.Fprintf(sb, "error: %v\n", err) + return + } + + if len(names) == 0 { + sb.WriteString("no rules\n") + return + } + + for _, name := range names { + owner := "" + if strings.HasPrefix(strings.ToLower(name), strings.ToLower(nbdns.NRPTKeyPrefix)) { + owner = " (netbird)" + } + fmt.Fprintf(sb, "%s%s\n", name, owner) + g.writeValues(sb, root+`\`+name, nil, " ") + } +} + +// writeEffectiveNRPTPolicies reports the table the resolver answers from, which +// the registry cannot show: a rule is written before it is loaded, and it keeps +// being enforced after its key is gone until the resolver reloads its policy. +func (g *BundleGenerator) writeEffectiveNRPTPolicies(sb *strings.Builder) { + writeSection(sb, "NRPT policy table in effect", nrptPolicyClass+"."+nrptPolicyMethod+" in "+nrptPolicyNamespace) + + entries, err := effectiveNRPTPolicies() + if err != nil { + fmt.Fprintf(sb, "error: %v\n", err) + return + } + + if len(entries) == 0 { + sb.WriteString("no policies\n") + return + } + + for _, entry := range entries { + fmt.Fprintf(sb, "%s\n", g.anonymizeValue("Namespace", entry.namespace)) + for _, value := range entry.values { + fmt.Fprintf(sb, " %s: %s\n", value.name, g.anonymizeValue(value.name, value.value)) + } + } +} + +// writeInterfaceDNS reports the DNS values of every interface that has any, so +// the netbird interface can be compared against the physical ones. The registry +// keys the values by GUID, so each is named from the adapter list; a GUID with +// no adapter is a leftover key of an interface that no longer exists. +func (g *BundleGenerator) writeInterfaceDNS(sb *strings.Builder, title, root string, names map[string]string) { + writeSection(sb, title, root) + + guids, err := subKeyNames(root) + if err != nil { + fmt.Fprintf(sb, "error: %v\n", err) + return + } + + var reported int + for _, guid := range guids { + var iface strings.Builder + g.writeValues(&iface, root+`\`+guid, interfaceDNSValues, " ") + if iface.Len() == 0 { + continue + } + + name, ok := names[strings.ToLower(guid)] + if !ok { + name = "no adapter with this GUID" + } + + reported++ + fmt.Fprintf(sb, "%s (%s)\n%s", guid, name, iface.String()) + } + + if reported == 0 { + sb.WriteString("no interface holds DNS values\n") + } +} + +// writeRegistryKey reports the values of a single key, without its subkeys. +func (g *BundleGenerator) writeRegistryKey(sb *strings.Builder, title, path string) { + writeSection(sb, title, path) + + var values strings.Builder + g.writeValues(&values, path, nil, "") + if values.Len() == 0 { + sb.WriteString("no values\n") + return + } + + sb.WriteString(values.String()) +} + +// writeValues renders the values of a key. A nil names list reports every +// value, otherwise only those named and present. +func (g *BundleGenerator) writeValues(sb *strings.Builder, path string, names []string, indent string) { + k, err := registry.OpenKey(registry.LOCAL_MACHINE, path, registry.QUERY_VALUE) + switch { + case errors.Is(err, registry.ErrNotExist), errors.Is(err, windows.ERROR_PATH_NOT_FOUND): + // an absent key is the normal state for the GPO store and for + // interfaces without DNS settings + log.Debugf("HKEY_LOCAL_MACHINE\\%s does not exist", path) + return + case err != nil: + fmt.Fprintf(sb, "%serror: open HKEY_LOCAL_MACHINE\\%s: %v\n", indent, path, err) + return + } + defer closeKey(k) + + if names == nil { + names, err = k.ReadValueNames(-1) + if err != nil { + fmt.Fprintf(sb, "%serror: read value names: %v\n", indent, err) + return + } + } + + for _, name := range names { + value, err := readRegistryValue(k, name) + switch { + case errors.Is(err, registry.ErrNotExist): + // the caller asks for a fixed set of values, most of which a + // given interface does not carry + continue + case err != nil: + // report rather than omit: a value that is there but cannot be + // read reads as unset otherwise + fmt.Fprintf(sb, "%s%s: error: %v\n", indent, name, err) + continue + } + + fmt.Fprintf(sb, "%s%s: %s\n", indent, name, g.anonymizeValue(name, value)) + } +} + +// anonymizeValue redacts a registry value according to what its name says it +// holds. Domains and addresses are handled per entry rather than by the string +// pass: the pass only replaces domains something else in the bundle already +// seeded, and its address regex would eat the digit labels of a reverse zone. +func (g *BundleGenerator) anonymizeValue(name, value string) string { + if !g.anonymize || value == "" { + return value + } + + switch { + case holdsDomains(name): + return joinValueEntries(splitValueEntries(value), g.anonymizeDomain) + case holdsAddresses(name): + return joinValueEntries(splitValueEntries(value), g.anonymizer.AnonymizeIPString) + default: + return g.anonymizer.AnonymizeString(value) + } +} + +// holdsDomains reports whether a value name holds domains: the domain list of +// an NRPT rule (Name) or of the policy table (Namespace), a search list, the +// DNS suffix values of the TCP/IP and policy keys, which all end in "Domain" +// (Domain, DhcpDomain, NV Domain, ICSDomain), and a proxy host name. +func holdsDomains(name string) bool { + lower := strings.ToLower(name) + return lower == "name" || lower == "namespace" || lower == "searchlist" || + strings.HasSuffix(lower, "domain") || strings.HasSuffix(lower, "proxyname") +} + +// holdsAddresses reports whether a value name holds DNS server addresses +// (NameServer, DhcpNameServer, GenericDNSServers, NameServers). +func holdsAddresses(name string) bool { + lower := strings.ToLower(name) + return strings.Contains(lower, "nameserver") || strings.Contains(lower, "dnsserver") +} + +// adapterNames maps adapter GUIDs, as the registry keys the interfaces, to the +// names an operator sees. +func adapterNames(adapters []*windows.IpAdapterAddresses) map[string]string { + names := make(map[string]string, len(adapters)) + for _, adapter := range adapters { + guid := windows.BytePtrToString(adapter.AdapterName) + names[strings.ToLower(guid)] = windows.UTF16PtrToString(adapter.FriendlyName) + } + return names +} + +// writeAdapterDNS reports the resolver configuration in effect per adapter, +// which is what the resolver uses for a name no NRPT rule matches. +func (g *BundleGenerator) writeAdapterDNS(sb *strings.Builder, adapters []*windows.IpAdapterAddresses, err error) { + writeSection(sb, "Adapter DNS configuration", "GetAdaptersAddresses") + + if err != nil { + fmt.Fprintf(sb, "error: %v\n", err) + return + } + + for _, adapter := range adapters { + name := windows.UTF16PtrToString(adapter.FriendlyName) + suffix := g.anonymizeDomain(windows.UTF16PtrToString(adapter.DnsSuffix)) + + fmt.Fprintf(sb, "%s (index %d, oper status %d)\n", name, adapter.IfIndex, adapter.OperStatus) + fmt.Fprintf(sb, " DNS suffix: %s\n", suffix) + + var servers []string + for server := adapter.FirstDnsServerAddress; server != nil; server = server.Next { + addr, ok := netip.AddrFromSlice(server.Address.IP()) + if !ok { + continue + } + + addr = addr.Unmap() + if g.anonymize { + addr = g.anonymizer.AnonymizeIP(addr) + } + servers = append(servers, addr.String()) + } + + fmt.Fprintf(sb, " DNS servers: %s\n", strings.Join(servers, ", ")) + } +} + +// anonymizeDomain anonymizes a single domain, keeping the leading dot an NRPT +// match domain carries. +func (g *BundleGenerator) anonymizeDomain(entry string) string { + if !g.anonymize { + return entry + } + + domain, dot := strings.CutPrefix(entry, ".") + if domain == "" { + return entry + } + + anonymized := g.anonymizer.AnonymizeDomain(domain) + if dot { + anonymized = "." + anonymized + } + return anonymized +} + +// splitValueEntries splits a registry value that holds a list. The separator +// differs per value: a REG_MULTI_SZ arrives joined with ", ", a SearchList is +// comma separated and a NameServer may use commas or spaces. +func splitValueEntries(value string) []string { + return strings.FieldsFunc(value, func(r rune) bool { + return r == ',' || r == ';' || r == ' ' || r == '\t' + }) +} + +func joinValueEntries(entries []string, anonymize func(string) string) string { + for i, entry := range entries { + entries[i] = anonymize(entry) + } + return strings.Join(entries, ", ") +} + +func writeSection(sb *strings.Builder, title, source string) { + fmt.Fprintf(sb, "\n%s\n%s\n%s\n", title, strings.Repeat("-", len(title)), source) +} + +func subKeyNames(root string) ([]string, error) { + k, err := registry.OpenKey(registry.LOCAL_MACHINE, root, registry.ENUMERATE_SUB_KEYS) + if err != nil { + return nil, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", root, err) + } + defer closeKey(k) + + names, err := k.ReadSubKeyNames(-1) + if err != nil { + return nil, fmt.Errorf("read subkey names: %w", err) + } + + return names, nil +} + +// readRegistryValue renders a value as text regardless of its type, so an +// unexpected type in a policy key still shows up instead of being dropped. +func readRegistryValue(k registry.Key, name string) (string, error) { + _, valueType, err := k.GetValue(name, nil) + if err != nil { + return "", fmt.Errorf("get value %s: %w", name, err) + } + + switch valueType { + case registry.SZ, registry.EXPAND_SZ: + value, _, err := k.GetStringValue(name) + if err != nil { + return "", fmt.Errorf("get string value %s: %w", name, err) + } + return value, nil + case registry.MULTI_SZ: + values, _, err := k.GetStringsValue(name) + if err != nil { + return "", fmt.Errorf("get strings value %s: %w", name, err) + } + return strings.Join(values, ", "), nil + case registry.DWORD, registry.QWORD: + value, _, err := k.GetIntegerValue(name) + if err != nil { + return "", fmt.Errorf("get integer value %s: %w", name, err) + } + return fmt.Sprintf("%d (0x%x)", value, value), nil + case registry.BINARY: + value, _, err := k.GetBinaryValue(name) + if err != nil { + return "", fmt.Errorf("get binary value %s: %w", name, err) + } + return hex.EncodeToString(value), nil + default: + return fmt.Sprintf("", valueType), nil + } +} + +// adapterAddresses returns the adapter list including DNS servers. The call +// reports the size it needs, so grow the buffer and retry until it fits. +func adapterAddresses() (adapters []*windows.IpAdapterAddresses, err error) { + // GetAdaptersAddresses is resolved on first use and panics when it is + // missing, so this reports it as an error and leaves the rest of the + // report intact. + defer func() { + if r := recover(); r != nil { + adapters, err = nil, fmt.Errorf("GetAdaptersAddresses: %v", r) + } + }() + + const flags = windows.GAA_FLAG_SKIP_ANYCAST | windows.GAA_FLAG_SKIP_MULTICAST + + size := uint32(15000) + for range 3 { + buf := make([]byte, size) + first := (*windows.IpAdapterAddresses)(unsafe.Pointer(&buf[0])) + + err := windows.GetAdaptersAddresses(windows.AF_UNSPEC, flags, 0, first, &size) + if errors.Is(err, windows.ERROR_BUFFER_OVERFLOW) { + continue + } + if err != nil { + return nil, fmt.Errorf("GetAdaptersAddresses: %w", err) + } + + for adapter := first; adapter != nil; adapter = adapter.Next { + adapters = append(adapters, adapter) + } + return adapters, nil + } + + return nil, fmt.Errorf("GetAdaptersAddresses: buffer kept growing") +} + +func closeKey(k registry.Key) { + if err := k.Close(); err != nil { + log.Debugf("close registry key: %v", err) + } +} diff --git a/client/internal/debug/debug_windows_test.go b/client/internal/debug/debug_windows_test.go new file mode 100644 index 000000000..47df3f6f9 --- /dev/null +++ b/client/internal/debug/debug_windows_test.go @@ -0,0 +1,146 @@ +//go:build windows + +package debug + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/anonymize" +) + +func newDNSValueGenerator(level anonymize.Level) *BundleGenerator { + anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(level) + + return &BundleGenerator{ + anonymize: true, + anonymizeLevel: level, + anonymizer: anonymizer, + } +} + +// TestAnonymizeValueByName covers the value kinds of the DNS registry keys. The +// names decide the treatment, because the string pass alone replaces only +// domains another part of the bundle already seeded. +func TestAnonymizeValueByName(t *testing.T) { + tests := []struct { + name string + valueName string + value string + assert func(t *testing.T, got string) + }{ + { + name: "NRPT match domains keep the leading dot", + valueName: "Name", + value: ".internal.example.com, .corp.example.org", + assert: func(t *testing.T, got string) { + t.Helper() + for _, entry := range strings.Split(got, ", ") { + assert.True(t, strings.HasPrefix(entry, "."), "entry %q should keep its leading dot", entry) + assert.NotContains(t, entry, "example", "entry %q should not keep the original domain", entry) + } + }, + }, + { + name: "any value name ending in Domain is treated as a domain", + valueName: "ICSDomain", + value: "mshome.net", + assert: func(t *testing.T, got string) { + t.Helper() + assert.NotContains(t, got, "mshome", "should anonymize a domain suffix value") + }, + }, + { + name: "search list is a comma separated domain list", + valueName: "SearchList", + value: "corp.example.com,branch.example.com", + assert: func(t *testing.T, got string) { + t.Helper() + assert.NotContains(t, got, "example", "should anonymize every search domain") + assert.Len(t, strings.Split(got, ", "), 2, "should keep both search domains") + }, + }, + { + name: "name servers are anonymized as addresses", + valueName: "DhcpNameServer", + value: "203.0.113.10 8.8.8.8", + assert: func(t *testing.T, got string) { + t.Helper() + assert.NotContains(t, got, "203.0.113.10", "should anonymize a public resolver address") + // well-known resolvers stay readable at every level + assert.Contains(t, got, "8.8.8.8", "should keep a well-known resolver address") + }, + }, + { + name: "opaque values are left to the string pass", + valueName: "DataBasePath", + value: `%SystemRoot%\System32\drivers\etc`, + assert: func(t *testing.T, got string) { + t.Helper() + assert.Equal(t, `%SystemRoot%\System32\drivers\etc`, got, "should not alter a path") + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + g := newDNSValueGenerator(anonymize.LevelDefault) + tc.assert(t, g.anonymizeValue(tc.valueName, tc.value)) + }) + } +} + +// TestParseNRPTPolicyTable parses the MOF text of the policy table out +// parameters, as the provider on a client with one NRPT rule renders it. +func TestParseNRPTPolicyTable(t *testing.T) { + const text = `[abstract] +class __PARAMETERS +{ + [Out, EmbeddedInstance("DnsClientPolicyConfiguration"): ToSubClass, ID(2): DisableOverride ToInstance] DnsClientPolicyConfiguration cmdletOutput[] = { +instance of DnsClientPolicyConfiguration +{ + DirectAccessProxyType = "NoProxy"; + DirectAccessQueryIPsecRequired = FALSE; + NameEncoding = "Utf8WithoutMapping"; + Namespace = ".0.100.in-addr.arpa"; +}, +instance of DnsClientPolicyConfiguration +{ + DirectAccessProxyType = "NoProxy"; + NameEncoding = "Utf8WithoutMapping"; + NameServers = {"100.0.255.254", "100.0.255.253"}; + Namespace = ".nb.internal"; +}}; + [in] boolean Effective; + [out] uint32 ReturnValue = 0; +}; +` + + entries := parseNRPTPolicyTable(text) + require.Len(t, entries, 2, "should parse both embedded instances") + + assert.Equal(t, ".0.100.in-addr.arpa", entries[0].namespace, "should read the namespace of the first instance") + assert.Equal(t, ".nb.internal", entries[1].namespace, "should read the namespace of the second instance") + + assert.Equal(t, []registryValue{ + {name: "DirectAccessProxyType", value: "NoProxy"}, + {name: "DirectAccessQueryIPsecRequired", value: "FALSE"}, + {name: "NameEncoding", value: "Utf8WithoutMapping"}, + }, entries[0].values, "should keep the remaining values in order") + + assert.Contains(t, entries[1].values, registryValue{name: "NameServers", value: "100.0.255.254, 100.0.255.253"}, + "should flatten a MOF array") + + for _, value := range entries[1].values { + assert.NotContains(t, value.name, "ReturnValue", "should not read the class level parameters as values") + } +} + +func TestParseNRPTPolicyTableEmpty(t *testing.T) { + assert.Empty(t, parseNRPTPolicyTable(""), "should parse no entries from empty text") + assert.Empty(t, parseNRPTPolicyTable("class __PARAMETERS\n{\n};\n"), "should parse no entries from a table with no instances") +} diff --git a/client/internal/debug/nrpt_windows.go b/client/internal/debug/nrpt_windows.go new file mode 100644 index 000000000..6b6e0e29a --- /dev/null +++ b/client/internal/debug/nrpt_windows.go @@ -0,0 +1,317 @@ +//go:build windows + +package debug + +import ( + "errors" + "fmt" + "runtime" + "strings" + "time" + + "github.com/go-ole/go-ole" + "github.com/go-ole/go-ole/oleutil" + log "github.com/sirupsen/logrus" +) + +const ( + // The NRPT policy table is reachable through the CIM class that backs + // Get-DnsClientNrptPolicy. Unlike the rules in the registry, the table is + // what the resolver currently has loaded, which is the only way to tell an + // applied rule from one that is merely written, in either direction. + nrptPolicyNamespace = `root\Microsoft\Windows\DNS` + nrptPolicyClass = "PS_DnsClientNrptPolicy" + nrptPolicyMethod = "Get" + + // The class has no instances, so the table comes from the out parameters + // of a static method call, rendered as MOF text: the embedded instances + // arrive as a safe array of objects, which cannot be read back through the + // COM bindings, and the text form carries all of them. + nrptPolicyInstanceKeyword = "instance of DnsClientPolicyConfiguration" + + nrptPolicyTimeout = 15 * time.Second +) + +// COM initialization results that leave the calling thread usable: S_FALSE for +// a thread this process already initialized, RPC_E_CHANGED_MODE for one that +// belongs to another apartment. +const ( + sFalse = 0x00000001 + rpcEChangedMode = 0x80010106 +) + +// nrptQueryInFlight admits one read of the policy table at a time. A provider +// that stops answering keeps its goroutine and the OS thread that goroutine +// pinned, so a later bundle reports that instead of pinning another one. +var nrptQueryInFlight = make(chan struct{}, 1) + +// nrptPolicyEntry is one namespace of the effective policy table, holding the +// values of an embedded DnsClientPolicyConfiguration instance in the order the +// provider reported them. +type nrptPolicyEntry struct { + namespace string + values []registryValue +} + +// registryValue is a name and its rendered value, shared by the registry and +// policy table readers so both anonymize by value name the same way. +type registryValue struct { + name string + value string +} + +// effectiveNRPTPolicies reads the effective NRPT table. The call is bounded +// because a WMI provider can block indefinitely and a debug bundle must not. +func effectiveNRPTPolicies() ([]nrptPolicyEntry, error) { + type result struct { + text string + err error + } + + select { + case nrptQueryInFlight <- struct{}{}: + default: + return nil, errors.New("an earlier read of the policy table has not returned") + } + + done := make(chan result, 1) + go func() { + // the slot is released here rather than by the caller, so a read that + // outlives the timeout holds it until the provider answers + defer func() { <-nrptQueryInFlight }() + + text, err := nrptPolicyTableText() + done <- result{text: text, err: err} + }() + + select { + case res := <-done: + if res.err != nil { + return nil, res.err + } + return parseNRPTPolicyTable(res.text), nil + case <-time.After(nrptPolicyTimeout): + return nil, errors.New("read of the policy table timed out") + } +} + +// nrptPolicyTableText calls the policy table method and returns the MOF text of +// its out parameters. +func nrptPolicyTableText() (text string, err error) { + // COM is per thread, and the collection is short lived, so the thread is + // pinned for the duration rather than initialized for the process. + runtime.LockOSThread() + defer runtime.UnlockOSThread() + + defer func() { + // The COM call chain is dynamically typed, so a provider that answers + // with an unexpected shape must not take the daemon down with it. + if r := recover(); r != nil { + err = fmt.Errorf("read NRPT policy table: %v", r) + } + }() + + owns, err := coInitialize() + if err != nil { + return "", err + } + if owns { + defer ole.CoUninitialize() + } + + locator, err := oleutil.CreateObject("WbemScripting.SWbemLocator") + if err != nil { + return "", fmt.Errorf("create WMI locator: %w", err) + } + defer locator.Release() + + dispatch, err := locator.QueryInterface(ole.IID_IDispatch) + if err != nil { + return "", fmt.Errorf("query WMI locator interface: %w", err) + } + defer dispatch.Release() + + service, err := dispatchCall(dispatch, "ConnectServer", nil, nrptPolicyNamespace) + if err != nil { + return "", fmt.Errorf("connect to %s: %w", nrptPolicyNamespace, err) + } + defer service.Release() + + inParams, err := spawnMethodInParams(service) + if err != nil { + return "", err + } + defer inParams.Release() + + // The effective table is the merge of the local and the group policy + // store, which is what the resolver answers from. + if _, err := oleutil.PutProperty(inParams, "Effective", true); err != nil { + return "", fmt.Errorf("set Effective parameter: %w", err) + } + + outParams, err := dispatchCall(service, "ExecMethod", nrptPolicyClass, nrptPolicyMethod, inParams) + if err != nil { + return "", fmt.Errorf("call %s.%s: %w", nrptPolicyClass, nrptPolicyMethod, err) + } + defer outParams.Release() + + textVariant, err := oleutil.CallMethod(outParams, "GetObjectText_") + if err != nil { + return "", fmt.Errorf("render policy table: %w", err) + } + defer func() { + if err := textVariant.Clear(); err != nil { + log.Debugf("clear policy table variant: %v", err) + } + }() + + return textVariant.ToString(), nil +} + +// spawnMethodInParams builds the in parameters instance the method needs. The +// provider rejects the call without one, even when every parameter is optional. +func spawnMethodInParams(service *ole.IDispatch) (*ole.IDispatch, error) { + class, err := dispatchCall(service, "Get", nrptPolicyClass) + if err != nil { + return nil, fmt.Errorf("get class %s: %w", nrptPolicyClass, err) + } + defer class.Release() + + methods, err := dispatchProperty(class, "Methods_") + if err != nil { + return nil, fmt.Errorf("get class methods: %w", err) + } + defer methods.Release() + + method, err := dispatchCall(methods, "Item", nrptPolicyMethod) + if err != nil { + return nil, fmt.Errorf("get method %s: %w", nrptPolicyMethod, err) + } + defer method.Release() + + params, err := dispatchProperty(method, "InParameters") + if err != nil { + return nil, fmt.Errorf("get method parameters: %w", err) + } + defer params.Release() + + inParams, err := dispatchCall(params, "SpawnInstance_") + if err != nil { + return nil, fmt.Errorf("spawn parameter instance: %w", err) + } + + return inParams, nil +} + +// parseNRPTPolicyTable pulls the embedded instances out of the MOF text. Each +// instance is a namespace of the table, with one name and value per line. +func parseNRPTPolicyTable(text string) []nrptPolicyEntry { + var entries []nrptPolicyEntry + var current *nrptPolicyEntry + + for _, line := range strings.Split(text, "\n") { + line = strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(line), ";")) + + switch { + case strings.HasPrefix(line, nrptPolicyInstanceKeyword): + entries = append(entries, nrptPolicyEntry{}) + current = &entries[len(entries)-1] + continue + case strings.HasPrefix(line, "}"): + // closes an instance, and the array with the last one, so the + // class level parameters that follow are not read as values + current = nil + continue + case current == nil, line == "{": + continue + } + + name, value, ok := strings.Cut(line, " = ") + if !ok { + continue + } + + value = unquoteMOFValue(value) + if name == "Namespace" { + current.namespace = value + continue + } + + current.values = append(current.values, registryValue{name: name, value: value}) + } + + return entries +} + +// unquoteMOFValue renders a MOF scalar or array as plain text: "a" becomes a, +// and {"a", "b"} becomes a, b. +func unquoteMOFValue(value string) string { + value = strings.TrimSpace(value) + + if inner, ok := strings.CutPrefix(value, "{"); ok { + value = strings.TrimSuffix(inner, "}") + + entries := strings.Split(value, ",") + for i, entry := range entries { + entries[i] = strings.Trim(strings.TrimSpace(entry), `"`) + } + return strings.Join(entries, ", ") + } + + return strings.Trim(value, `"`) +} + +// coInitialize prepares the calling thread for COM and reports whether this +// call owns the initialization, which decides whether it may be balanced with +// CoUninitialize. S_FALSE took a reference on a thread this process had already +// initialized and so has to be released, while RPC_E_CHANGED_MODE took none: +// the thread belongs to another apartment, which is usable but is not ours to +// uninitialize. +func coInitialize() (bool, error) { + err := ole.CoInitializeEx(0, ole.COINIT_MULTITHREADED) + if err == nil { + return true, nil + } + + var oleErr *ole.OleError + if errors.As(err, &oleErr) { + switch oleErr.Code() { + case sFalse: + return true, nil + case rpcEChangedMode: + return false, nil + } + } + + return false, fmt.Errorf("initialize COM: %w", err) +} + +// dispatchCall calls a COM method that returns an object. +func dispatchCall(dispatch *ole.IDispatch, method string, params ...any) (*ole.IDispatch, error) { + variant, err := oleutil.CallMethod(dispatch, method, params...) + if err != nil { + return nil, err + } + + object := variant.ToIDispatch() + if object == nil { + return nil, fmt.Errorf("%s returned no object", method) + } + + return object, nil +} + +// dispatchProperty reads a COM property that holds an object. +func dispatchProperty(dispatch *ole.IDispatch, property string) (*ole.IDispatch, error) { + variant, err := oleutil.GetProperty(dispatch, property) + if err != nil { + return nil, err + } + + object := variant.ToIDispatch() + if object == nil { + return nil, fmt.Errorf("property %s holds no object", property) + } + + return object, nil +} diff --git a/client/internal/debug/uilog_test.go b/client/internal/debug/uilog_test.go new file mode 100644 index 000000000..103e98c6f --- /dev/null +++ b/client/internal/debug/uilog_test.go @@ -0,0 +1,64 @@ +package debug + +import ( + "archive/zip" + "bytes" + "fmt" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/configs" +) + +// bundleEntries generates a bundle with the given generator and returns the +// set of entry names in the resulting archive. +func bundleEntries(t *testing.T, g *BundleGenerator) map[string]struct{} { + t.Helper() + + path, err := g.Generate() + require.NoError(t, err) + t.Cleanup(func() { _ = os.Remove(path) }) + + data, err := os.ReadFile(path) + require.NoError(t, err) + + zr, err := zip.NewReader(bytes.NewReader(data), int64(len(data))) + require.NoError(t, err) + + names := make(map[string]struct{}, len(zr.File)) + for _, f := range zr.File { + names[f.Name] = struct{}{} + } + return names +} + +func TestBundleIncludesUILogWhenOpenerAllows(t *testing.T) { + path := filepath.Join(t.TempDir(), configs.UILogFile) + require.NoError(t, os.WriteFile(path, []byte("gui log"), 0600)) + + g := NewBundleGenerator(GeneratorDependencies{ + UILogPath: path, + UILogOpener: openLogFile, + }, BundleConfig{}) + + require.Contains(t, bundleEntries(t, g), configs.UILogFile) +} + +// A UILogOpener that refuses (as the ownership check does for a foreign file) +// keeps the UI log out of the bundle without failing bundle generation. +func TestBundleExcludesUILogWhenOpenerRefuses(t *testing.T) { + path := filepath.Join(t.TempDir(), configs.UILogFile) + require.NoError(t, os.WriteFile(path, []byte("secret"), 0600)) + + g := NewBundleGenerator(GeneratorDependencies{ + UILogPath: path, + UILogOpener: func(string) (*os.File, error) { + return nil, fmt.Errorf("not owned by the caller") + }, + }, BundleConfig{}) + + require.NotContains(t, bundleEntries(t, g), configs.UILogFile) +} diff --git a/client/internal/debug/upload.go b/client/internal/debug/upload.go index cdf52409d..88fde6d6f 100644 --- a/client/internal/debug/upload.go +++ b/client/internal/debug/upload.go @@ -3,10 +3,12 @@ package debug import ( "context" "crypto/sha256" + "crypto/tls" "encoding/json" "fmt" "io" "net/http" + neturl "net/url" "os" "github.com/netbirdio/netbird/upload-server/types" @@ -14,20 +16,80 @@ import ( const maxBundleUploadSize = 50 * 1024 * 1024 -func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string) (key string, err error) { - response, err := getUploadURL(ctx, url, managementURL) +// requireHTTPS refuses any URL the daemon would fetch or upload to that is not +// https. The daemon runs as root and the bundle carries its logs and state, so a +// plaintext hop is a place to intercept the bundle or the presigned redirect. +// The server-side gate already enforces this for the desktop path; this also +// covers the mobile and job-runner callers that reach this package directly. +// Skipped when the caller opted into an insecure upload (self-hosted server). +func requireHTTPS(what, rawURL string) error { + parsed, err := neturl.Parse(rawURL) + if err != nil { + return fmt.Errorf("parse %s: %w", what, err) + } + if parsed.Scheme != "https" { + return fmt.Errorf("%s must use https, got scheme %q", what, parsed.Scheme) + } + return nil +} + +// uploadClient returns the HTTP client for the upload requests. The default +// client verifies TLS and refuses a redirect that would downgrade to a non-https +// hop, so a bundle can never leave over http after an https start. The insecure +// variant accepts http and untrusted certificates, and is only reachable for a +// privileged caller that passed --upload-bundle-insecure (see +// requirePrivilegeForUploadURL). +func uploadClient(insecure bool) *http.Client { + if !insecure { + return &http.Client{CheckRedirect: rejectInsecureRedirect} + } + return &http.Client{ + Transport: &http.Transport{ + //nolint:gosec // opt-in, privileged, self-hosted upload servers + TLSClientConfig: &tls.Config{InsecureSkipVerify: true, MinVersion: tls.VersionTLS12}, + }, + } +} + +// rejectInsecureRedirect refuses a redirect to a non-https target and keeps the +// standard library's 10-hop limit that a custom CheckRedirect would otherwise +// disable. +func rejectInsecureRedirect(req *http.Request, via []*http.Request) error { + if req.URL.Scheme != "https" { + return fmt.Errorf("refusing redirect to non-https URL %s", req.URL.Redacted()) + } + if len(via) >= 10 { + return fmt.Errorf("stopped after 10 redirects") + } + return nil +} + +func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string, insecure bool) (key string, err error) { + if !insecure { + if err := requireHTTPS("upload service URL", url); err != nil { + return "", err + } + } + + response, err := getUploadURL(ctx, url, managementURL, insecure) if err != nil { return "", err } - err = upload(ctx, filePath, response) + if !insecure { + if err := requireHTTPS("upload URL from service", response.URL); err != nil { + return "", err + } + } + + err = upload(ctx, filePath, response, insecure) if err != nil { return "", err } return response.Key, nil } -func upload(ctx context.Context, filePath string, response *types.GetURLResponse) error { +func upload(ctx context.Context, filePath string, response *types.GetURLResponse, insecure bool) error { fileData, err := os.Open(filePath) if err != nil { return fmt.Errorf("open file: %w", err) @@ -52,7 +114,7 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse req.ContentLength = stat.Size() req.Header.Set("Content-Type", "application/octet-stream") - putResp, err := http.DefaultClient.Do(req) + putResp, err := uploadClient(insecure).Do(req) if err != nil { return fmt.Errorf("upload failed: %v", err) } @@ -65,16 +127,23 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse return nil } -func getUploadURL(ctx context.Context, url string, managementURL string) (*types.GetURLResponse, error) { - id := getURLHash(managementURL) - getReq, err := http.NewRequestWithContext(ctx, "GET", url+"?id="+id, nil) +func getUploadURL(ctx context.Context, serviceURL string, managementURL string, insecure bool) (*types.GetURLResponse, error) { + parsed, err := neturl.Parse(serviceURL) + if err != nil { + return nil, fmt.Errorf("parse upload service URL: %w", err) + } + q := parsed.Query() + q.Set("id", getURLHash(managementURL)) + parsed.RawQuery = q.Encode() + + getReq, err := http.NewRequestWithContext(ctx, "GET", parsed.String(), nil) if err != nil { return nil, fmt.Errorf("create GET request: %w", err) } getReq.Header.Set(types.ClientHeader, types.ClientHeaderValue) - resp, err := http.DefaultClient.Do(getReq) + resp, err := uploadClient(insecure).Do(getReq) if err != nil { return nil, fmt.Errorf("get presigned URL: %w", err) } diff --git a/client/internal/debug/upload_test.go b/client/internal/debug/upload_test.go index f224b8d3f..f3927cb81 100644 --- a/client/internal/debug/upload_test.go +++ b/client/internal/debug/upload_test.go @@ -5,6 +5,7 @@ import ( "errors" "net" "net/http" + "net/http/httptest" "os" "path/filepath" "testing" @@ -43,7 +44,7 @@ func TestUpload(t *testing.T) { fileContent := []byte("test file content") err := os.WriteFile(file, fileContent, 0640) require.NoError(t, err) - key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file) + key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file, true) require.NoError(t, err) id := getURLHash(testURL) require.Contains(t, key, id+"/") @@ -79,3 +80,47 @@ func waitForServer(t *testing.T, addr string) { } t.Fatalf("server did not start listening on %s in time", addr) } + +func TestRequireHTTPS(t *testing.T) { + require.NoError(t, requireHTTPS("upload URL", "https://upload.example/path")) + require.Error(t, requireHTTPS("upload URL", "http://upload.example/path")) + require.Error(t, requireHTTPS("upload URL", "ftp://upload.example/path")) + require.Error(t, requireHTTPS("upload URL", "://malformed")) +} + +func TestRejectInsecureRedirect(t *testing.T) { + httpsReq, err := http.NewRequest(http.MethodGet, "https://a.example/", nil) + require.NoError(t, err) + require.NoError(t, rejectInsecureRedirect(httpsReq, nil), "https redirect target must be allowed") + + httpReq, err := http.NewRequest(http.MethodGet, "http://a.example/", nil) + require.NoError(t, err) + require.Error(t, rejectInsecureRedirect(httpReq, nil), "http redirect target must be refused") + + require.Error(t, rejectInsecureRedirect(httpsReq, make([]*http.Request, 10)), "the 10-redirect limit must be enforced") +} + +// The secure client refuses to follow an https response that redirects to http, +// so a bundle can't be downgraded onto plaintext mid-flight. +func TestUploadClientRefusesHTTPSToHTTPRedirect(t *testing.T) { + plain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(plain.Close) + + secure := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, plain.URL, http.StatusFound) + })) + t.Cleanup(secure.Close) + + client := uploadClient(false) + // Trust the test server's cert without disabling verification globally. + client.Transport = secure.Client().Transport + + resp, err := client.Get(secure.URL) + if resp != nil { + _ = resp.Body.Close() + } + require.Error(t, err, "redirect from https to http must be refused") + require.Contains(t, err.Error(), "non-https") +} diff --git a/client/internal/debug/wgshow.go b/client/internal/debug/wgshow.go index 1e8a8a6cc..ee24902e6 100644 --- a/client/internal/debug/wgshow.go +++ b/client/internal/debug/wgshow.go @@ -35,14 +35,14 @@ func (g *BundleGenerator) toWGShowFormat(s *configurer.Stats) string { var sb strings.Builder sb.WriteString(fmt.Sprintf("interface: %s\n", s.DeviceName)) - sb.WriteString(fmt.Sprintf(" public key: %s\n", s.PublicKey)) + sb.WriteString(fmt.Sprintf(" public key: %s\n", g.anonymizer.AnonymizeWGKey(s.PublicKey))) sb.WriteString(fmt.Sprintf(" listen port: %d\n", s.ListenPort)) if s.FWMark != 0 { sb.WriteString(fmt.Sprintf(" fwmark: %#x\n", s.FWMark)) } for _, peer := range s.Peers { - sb.WriteString(fmt.Sprintf("\npeer: %s\n", peer.PublicKey)) + sb.WriteString(fmt.Sprintf("\npeer: %s\n", g.anonymizer.AnonymizeWGKey(peer.PublicKey))) if peer.Endpoint.IP != nil { if g.anonymize { anonEndpoint := g.anonymizer.AnonymizeUDPAddr(peer.Endpoint) @@ -54,7 +54,11 @@ func (g *BundleGenerator) toWGShowFormat(s *configurer.Stats) string { if len(peer.AllowedIPs) > 0 { var ipStrings []string for _, ipnet := range peer.AllowedIPs { - ipStrings = append(ipStrings, ipnet.String()) + ipStr := ipnet.String() + if g.anonymize { + ipStr = g.anonymizer.AnonymizeIPString(ipStr) + } + ipStrings = append(ipStrings, ipStr) } sb.WriteString(fmt.Sprintf(" allowed ips: %s\n", strings.Join(ipStrings, ", "))) } diff --git a/client/internal/dns/host_darwin.go b/client/internal/dns/host_darwin.go index 0f4eb6bf8..81029752e 100644 --- a/client/internal/dns/host_darwin.go +++ b/client/internal/dns/host_darwin.go @@ -267,18 +267,38 @@ func (s *systemConfigurator) getSystemDNSSettings() (SystemDNSSettings, error) { return SystemDNSSettings{}, fmt.Errorf("sending the command: %w", err) } - var dnsSettings SystemDNSSettings + dnsSettings, serverAddresses, err := parseSystemDNSSettings(b) + if err != nil { + return dnsSettings, err + } + + s.mu.Lock() + s.origNameservers = serverAddresses + s.mu.Unlock() + + return dnsSettings, nil +} + +// parseSystemDNSSettings parses the output of `scutil show State:/Network/Service//DNS`. +// Lines that don't match the expected "index : value" shape are skipped: hosts with unusual +// network services (e.g. orphaned hardware ports) can produce entries without a value. +func parseSystemDNSSettings(out []byte) (SystemDNSSettings, []netip.Addr, error) { + // port is not exposed by scutil, default to 53 + dnsSettings := SystemDNSSettings{ServerPort: DefaultPort} var serverAddresses []netip.Addr inSearchDomainsArray := false inServerAddressesArray := false - scanner := bufio.NewScanner(bytes.NewReader(b)) + scanner := bufio.NewScanner(bytes.NewReader(out)) for scanner.Scan() { line := strings.TrimSpace(scanner.Text()) switch { case strings.HasPrefix(line, "DomainName :"): - domainName := strings.TrimSpace(strings.Split(line, ":")[1]) - dnsSettings.Domains = append(dnsSettings.Domains, domainName) + domainName := strings.TrimSpace(strings.TrimPrefix(line, "DomainName :")) + if domainName != "" { + dnsSettings.Domains = append(dnsSettings.Domains, domainName) + } + continue case line == "SearchDomains : {": inSearchDomainsArray = true continue @@ -288,36 +308,45 @@ func (s *systemConfigurator) getSystemDNSSettings() (SystemDNSSettings, error) { case line == "}": inSearchDomainsArray = false inServerAddressesArray = false + continue + } + + if !inSearchDomainsArray && !inServerAddressesArray { + continue + } + + parts := strings.SplitN(line, " : ", 2) + if len(parts) != 2 { + log.Debugf("skipping unexpected scutil DNS line %q", line) + continue + } + value := strings.TrimSpace(parts[1]) + if value == "" { + continue } if inSearchDomainsArray { - searchDomain := strings.Split(line, " : ")[1] - dnsSettings.Domains = append(dnsSettings.Domains, searchDomain) - } else if inServerAddressesArray { - address := strings.Split(line, " : ")[1] - if ip, err := netip.ParseAddr(address); err == nil && !ip.IsUnspecified() { - ip = ip.Unmap() - serverAddresses = append(serverAddresses, ip) - // Prefer the first IPv4 server as ServerIP since our DNS listener is IPv4. - if !dnsSettings.ServerIP.IsValid() && ip.Is4() { - dnsSettings.ServerIP = ip - } - } + dnsSettings.Domains = append(dnsSettings.Domains, value) + continue + } + + ip, err := netip.ParseAddr(value) + if err != nil || ip.IsUnspecified() { + continue + } + ip = ip.Unmap() + serverAddresses = append(serverAddresses, ip) + // Prefer the first IPv4 server as ServerIP since our DNS listener is IPv4. + if !dnsSettings.ServerIP.IsValid() && ip.Is4() { + dnsSettings.ServerIP = ip } } if err := scanner.Err(); err != nil { - return dnsSettings, err + return dnsSettings, serverAddresses, err } - // default to 53 port - dnsSettings.ServerPort = DefaultPort - - s.mu.Lock() - s.origNameservers = serverAddresses - s.mu.Unlock() - - return dnsSettings, nil + return dnsSettings, serverAddresses, nil } func (s *systemConfigurator) getOriginalNameservers() []netip.Addr { @@ -435,11 +464,15 @@ func (s *systemConfigurator) getPrimaryService() (string, string, error) { router := "" for scanner.Scan() { text := scanner.Text() + parts := strings.SplitN(text, ":", 2) + if len(parts) != 2 { + continue + } if strings.Contains(text, "PrimaryService") { - primaryService = strings.TrimSpace(strings.Split(text, ":")[1]) + primaryService = strings.TrimSpace(parts[1]) } if strings.Contains(text, "Router") { - router = strings.TrimSpace(strings.Split(text, ":")[1]) + router = strings.TrimSpace(parts[1]) } } if err := scanner.Err(); err != nil && err != io.EOF { diff --git a/client/internal/dns/host_darwin_test.go b/client/internal/dns/host_darwin_test.go index 94d020c39..bee691c71 100644 --- a/client/internal/dns/host_darwin_test.go +++ b/client/internal/dns/host_darwin_test.go @@ -328,6 +328,120 @@ func removeTestDNSKey(key string) error { return err } +func TestParseSystemDNSSettings(t *testing.T) { + tests := []struct { + name string + output string + expectedDomains []string + expectedServers []netip.Addr + expectedIP netip.Addr + }{ + { + name: "well_formed", + output: ` { + DomainName : example.com + SearchDomains : { + 0 : example.com + 1 : corp.example.com + } + ServerAddresses : { + 0 : 192.168.1.1 + 1 : fd00::53 + } +} +`, + expectedDomains: []string{"example.com", "example.com", "corp.example.com"}, + expectedServers: []netip.Addr{netip.MustParseAddr("192.168.1.1"), netip.MustParseAddr("fd00::53")}, + expectedIP: netip.MustParseAddr("192.168.1.1"), + }, + { + // entries without a value after the separator used to panic with + // "index out of range [1] with length 1" + name: "malformed_array_entries_skipped", + output: ` { + SearchDomains : { + 0 : + (null) + + 1 : corp.example.com + } + ServerAddresses : { + 0 : + 1 : 192.168.1.1 + } +} +`, + expectedDomains: []string{"corp.example.com"}, + expectedServers: []netip.Addr{netip.MustParseAddr("192.168.1.1")}, + expectedIP: netip.MustParseAddr("192.168.1.1"), + }, + { + name: "domain_name_without_value_skipped", + output: ` { + DomainName : + ServerAddresses : { + 0 : 192.168.1.1 + } +} +`, + expectedServers: []netip.Addr{netip.MustParseAddr("192.168.1.1")}, + expectedIP: netip.MustParseAddr("192.168.1.1"), + }, + { + name: "ipv6_first_prefers_ipv4_server_ip", + output: ` { + ServerAddresses : { + 0 : fd00::53 + 1 : 192.168.1.1 + } +} +`, + expectedServers: []netip.Addr{netip.MustParseAddr("fd00::53"), netip.MustParseAddr("192.168.1.1")}, + expectedIP: netip.MustParseAddr("192.168.1.1"), + }, + { + name: "invalid_and_unspecified_addresses_skipped", + output: ` { + ServerAddresses : { + 0 : (null) + 1 : 0.0.0.0 + 2 : 192.168.1.1 + } +} +`, + expectedServers: []netip.Addr{netip.MustParseAddr("192.168.1.1")}, + expectedIP: netip.MustParseAddr("192.168.1.1"), + }, + { + name: "v4_mapped_address_unmapped", + output: ` { + ServerAddresses : { + 0 : ::ffff:192.168.1.1 + } +} +`, + expectedServers: []netip.Addr{netip.MustParseAddr("192.168.1.1")}, + expectedIP: netip.MustParseAddr("192.168.1.1"), + }, + { + name: "empty_output", + output: "", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + settings, servers, err := parseSystemDNSSettings([]byte(tc.output)) + require.NoError(t, err, "parsing should not fail") + + assert.Equal(t, tc.expectedDomains, settings.Domains, "domains should match") + assert.Equal(t, tc.expectedServers, servers, "server addresses should match") + assert.Equal(t, tc.expectedIP, settings.ServerIP, "server IP should match") + assert.Equal(t, DefaultPort, settings.ServerPort, "server port should default to 53") + }) + } +} + func TestGetOriginalNameservers(t *testing.T) { configurator := &systemConfigurator{ createdKeys: make(map[string]struct{}), diff --git a/client/internal/dns/host_windows.go b/client/internal/dns/host_windows.go index 4f6ece532..53380b2aa 100644 --- a/client/internal/dns/host_windows.go +++ b/client/internal/dns/host_windows.go @@ -31,10 +31,30 @@ var ( dnsFlushResolverCacheFn = dnsapi.NewProc("DnsFlushResolverCache") ) +// Registry locations of the host DNS configuration this package programs, +// exported so a diagnostic reader reports the same locations that are written. const ( - dnsPolicyConfigMatchPath = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters\DnsPolicyConfig\NetBird-Match` - gpoDnsPolicyRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient\DnsPolicyConfig` - gpoDnsPolicyConfigMatchPath = gpoDnsPolicyRoot + `\NetBird-Match` + // NRPTKeyPrefix starts the name of every NRPT rule key this client creates. + // Older versions used different layouts under the same prefix: a single + // unsuffixed key, then one key per domain, now one key per batch of domains. + NRPTKeyPrefix = "NetBird-Match" + + // DNSPolicyConfigRoot holds the NRPT rules of the local policy store. + DNSPolicyConfigRoot = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters\DnsPolicyConfig` + + // GPODNSPolicyConfigRoot holds the NRPT rules of the group policy store, + // which takes precedence over the local one when it is present. + GPODNSPolicyConfigRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient\DnsPolicyConfig` + + // InterfaceConfigPath and InterfaceConfigPathV6 hold the per-interface DNS + // settings, keyed by interface GUID, in separate hives per address family. + InterfaceConfigPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces` + InterfaceConfigPathV6 = `SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces` +) + +const ( + dnsPolicyConfigMatchPath = DNSPolicyConfigRoot + `\` + NRPTKeyPrefix + gpoDnsPolicyConfigMatchPath = GPODNSPolicyConfigRoot + `\` + NRPTKeyPrefix dnsPolicyConfigVersionKey = "Version" dnsPolicyConfigVersionValue = 2 @@ -45,8 +65,6 @@ const ( nrptMaxDomainsPerRule = 50 - interfaceConfigPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces` - interfaceConfigPathV6 = `SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces` interfaceConfigNameServerKey = "NameServer" interfaceConfigDhcpNameSrvKey = "DhcpNameServer" interfaceConfigSearchListKey = "SearchList" @@ -73,7 +91,6 @@ type registryConfigurator struct { guid string routingAll bool gpo bool - nrptEntryCount int origNameservers []netip.Addr } @@ -84,7 +101,7 @@ func newHostManager(wgInterface WGIface) (*registryConfigurator, error) { } var useGPO bool - k, err := registry.OpenKey(registry.LOCAL_MACHINE, gpoDnsPolicyRoot, registry.QUERY_VALUE) + 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 { @@ -123,7 +140,7 @@ func (r *registryConfigurator) captureOriginalNameservers() ([]netip.Addr, error seen := make(map[netip.Addr]struct{}) var out []netip.Addr var merr *multierror.Error - for _, root := range []string{interfaceConfigPath, interfaceConfigPathV6} { + for _, root := range []string{InterfaceConfigPath, InterfaceConfigPathV6} { addrs, err := r.captureFromTcpipRoot(root) if err != nil { merr = multierror.Append(merr, fmt.Errorf("%s: %w", root, err)) @@ -306,14 +323,9 @@ func (r *registryConfigurator) applyDNSConfig(config HostDNSConfig, stateManager } if len(matchDomains) != 0 { - count, err := r.addDNSMatchPolicy(matchDomains, config.ServerIP) - // Update count even on error to ensure cleanup covers partially created rules - r.nrptEntryCount = count - if err != nil { + if err := r.addDNSMatchPolicy(matchDomains, config.ServerIP); err != nil { return fmt.Errorf("add dns match policy: %w", err) } - } else { - r.nrptEntryCount = 0 } r.updateState(stateManager) @@ -329,9 +341,8 @@ func (r *registryConfigurator) applyDNSConfig(config HostDNSConfig, stateManager func (r *registryConfigurator) updateState(stateManager *statemanager.Manager) { if err := stateManager.UpdateState(&ShutdownState{ - Guid: r.guid, - GPO: r.gpo, - NRPTEntryCount: r.nrptEntryCount, + Guid: r.guid, + GPO: r.gpo, }); err != nil { log.Errorf("failed to update shutdown state: %s", err) } @@ -346,7 +357,7 @@ func (r *registryConfigurator) addDNSSetupForAll(ip netip.Addr) error { return nil } -func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr) (int, error) { +func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr) error { // if the gpo key is present, we need to put our DNS settings there, otherwise our config might be ignored // see https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-gpnrpt/8cc31cb9-20cb-4140-9e85-3e08703b4745 @@ -363,19 +374,17 @@ func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr gpoPath := fmt.Sprintf("%s-%d", gpoDnsPolicyConfigMatchPath, ruleIndex) if err := r.configureDNSPolicy(localPath, batchDomains, ip); err != nil { - return ruleIndex, fmt.Errorf("configure DNS Local policy for rule %d: %w", ruleIndex, err) + return fmt.Errorf("configure DNS Local policy for rule %d: %w", ruleIndex, err) } - // Increment immediately so the caller's cleanup path knows about this rule - ruleIndex++ - if r.gpo { if err := r.configureDNSPolicy(gpoPath, batchDomains, ip); err != nil { - return ruleIndex, fmt.Errorf("configure gpo DNS policy for rule %d: %w", ruleIndex-1, err) + return fmt.Errorf("configure gpo DNS policy for rule %d: %w", ruleIndex, err) } } - log.Debugf("added NRPT rule %d with %d domains", ruleIndex-1, len(batchDomains)) + log.Debugf("added NRPT rule %d with %d domains", ruleIndex, len(batchDomains)) + ruleIndex++ } if r.gpo { @@ -385,7 +394,7 @@ func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr } log.Infof("added %d NRPT rules for %d domains", ruleIndex, len(domains)) - return ruleIndex, nil + return nil } func (r *registryConfigurator) configureDNSPolicy(policyPath string, domains []string, ip netip.Addr) error { @@ -450,7 +459,7 @@ func (r *registryConfigurator) flushDNSCache() { ret, _, err := dnsFlushResolverCacheFn.Call() if ret == 0 { - if err != nil && !errors.Is(err, syscall.Errno(0)) { + if !errors.Is(err, syscall.Errno(0)) { log.Errorf("DnsFlushResolverCache failed: %v", err) return } @@ -496,7 +505,7 @@ func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey st } func (r *registryConfigurator) getInterfaceRegistryKey() (registry.Key, error) { - regKeyPath := interfaceConfigPath + "\\" + r.guid + regKeyPath := InterfaceConfigPath + "\\" + r.guid regKey, err := registry.OpenKey(registry.LOCAL_MACHINE, regKeyPath, registry.SET_VALUE) if err != nil { return regKey, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", regKeyPath, err) @@ -518,28 +527,28 @@ func (r *registryConfigurator) restoreHostDNS() error { return nil } +// removeDNSMatchPolicies deletes every NRPT rule this client may have created, +// from the local and the GPO policy store. The rules are found by enumerating +// the registry, the only authoritative record of what was written. Cleanup must +// not depend on a rule count: the in-memory one is scoped to a single +// registryConfigurator and the persisted one is deleted on every clean +// disconnect, and a rule left behind keeps resolving names over an interface +// that is gone, until reboot discards the volatile key. func (r *registryConfigurator) removeDNSMatchPolicies() error { var merr *multierror.Error - // Try to remove the base entries (for backward compatibility) - if err := removeRegistryKeyFromDNSPolicyConfig(dnsPolicyConfigMatchPath); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove local base entry: %w", err)) - } - - if err := removeRegistryKeyFromDNSPolicyConfig(gpoDnsPolicyConfigMatchPath); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove GPO base entry: %w", err)) - } - - for i := 0; i < r.nrptEntryCount; i++ { - localPath := fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i) - gpoPath := fmt.Sprintf("%s-%d", gpoDnsPolicyConfigMatchPath, i) - - if err := removeRegistryKeyFromDNSPolicyConfig(localPath); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove local entry %d: %w", i, err)) + for _, root := range []string{DNSPolicyConfigRoot, GPODNSPolicyConfigRoot} { + names, err := listNRPTRuleKeys(root) + if err != nil { + merr = multierror.Append(merr, fmt.Errorf("list rule keys under %s: %w", root, err)) + continue } - if err := removeRegistryKeyFromDNSPolicyConfig(gpoPath); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove GPO entry %d: %w", i, err)) + for _, name := range names { + path := root + `\` + name + if err := removeRegistryKeyFromDNSPolicyConfig(path); err != nil { + merr = multierror.Append(merr, fmt.Errorf("remove entry %s: %w", path, err)) + } } } @@ -554,6 +563,39 @@ func (r *registryConfigurator) restoreUncleanShutdownDNS() error { return r.restoreHostDNS() } +// 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. +func listNRPTRuleKeys(root string) ([]string, error) { + k, err := registry.OpenKey(registry.LOCAL_MACHINE, root, registry.ENUMERATE_SUB_KEYS) + switch { + case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND): + // the GPO store is absent on a machine without DNS client policy + log.Debugf("HKEY_LOCAL_MACHINE\\%s does not exist", root) + return nil, nil + case err != nil: + // any other failure has to reach the caller: reporting no rules would + // report a successful cleanup while leaving the rules in place + return nil, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", root, err) + } + defer closer(k) + + names, err := k.ReadSubKeyNames(-1) + if err != nil { + return nil, fmt.Errorf("read subkey names: %w", err) + } + + var ruleKeys []string + for _, name := range names { + // registry key names are case insensitive + if strings.HasPrefix(strings.ToLower(name), strings.ToLower(NRPTKeyPrefix)) { + ruleKeys = append(ruleKeys, name) + } + } + + return ruleKeys, nil +} + func removeRegistryKeyFromDNSPolicyConfig(regKeyPath string) error { k, err := registry.OpenKey(registry.LOCAL_MACHINE, regKeyPath, registry.QUERY_VALUE) if err != nil { @@ -585,7 +627,7 @@ func refreshGroupPolicy() error { ) if ret == 0 { - if err != nil && !errors.Is(err, syscall.Errno(0)) { + if !errors.Is(err, syscall.Errno(0)) { return fmt.Errorf("RefreshPolicyEx failed: %w", err) } return fmt.Errorf("RefreshPolicyEx failed") diff --git a/client/internal/dns/host_windows_test.go b/client/internal/dns/host_windows_test.go index 3cd2b1bd5..861613c95 100644 --- a/client/internal/dns/host_windows_test.go +++ b/client/internal/dns/host_windows_test.go @@ -25,7 +25,7 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) { // Create a test interface registry key so updateSearchDomains doesn't fail testGUID := "{12345678-1234-1234-1234-123456789ABC}" - interfacePath := `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\` + testGUID + 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() @@ -56,7 +56,7 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) { require.NoError(t, err) // Verify 3 NRPT rules exist - assert.Equal(t, 3, cfg.nrptEntryCount, "Should create 3 NRPT rules for 125 domains") + assert.Equal(t, 3, countNRPTRuleKeys(t), "Should create 3 NRPT rules for 125 domains") for i := 0; i < 3; i++ { exists, err := registryKeyExists(fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i)) require.NoError(t, err) @@ -81,7 +81,7 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) { require.NoError(t, err) // Verify first 2 NRPT rules exist - assert.Equal(t, 2, cfg.nrptEntryCount, "Should create 2 NRPT rules for 75 domains") + assert.Equal(t, 2, countNRPTRuleKeys(t), "Should create 2 NRPT rules for 75 domains") for i := 0; i < 2; i++ { exists, err := registryKeyExists(fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i)) require.NoError(t, err) @@ -106,9 +106,65 @@ func registryKeyExists(path string) (bool, error) { return true, nil } +// TestNRPTCleanupWithoutRuleCount verifies that rules written by a previous run +// are removed by a configurator that has no record of how many there are: an +// unclean exit loses the in-memory count and a clean disconnect deletes the +// persisted one, so cleanup cannot depend on either. +func TestNRPTCleanupWithoutRuleCount(t *testing.T) { + if testing.Short() { + t.Skip("skipping registry integration test in short mode") + } + + defer cleanupRegistryKeys(t) + cleanupRegistryKeys(t) + + testIP := netip.MustParseAddr("100.64.0.1") + + // 75 domains produce two indexed rules, as the current layout does + domains := make([]string, 75) + for i := range domains { + domains[i] = fmt.Sprintf(".domain%d.com", i+1) + } + + previousRun := ®istryConfigurator{} + require.NoError(t, previousRun.addDNSMatchPolicy(domains, testIP)) + + // the unsuffixed key an older version would have written + require.NoError(t, previousRun.configureDNSPolicy(dnsPolicyConfigMatchPath, []string{".legacy.example.com"}, testIP)) + + // a policy owned by someone else, which cleanup must not touch + foreignPath := DNSPolicyConfigRoot + `\DnsPolicyConfigTestForeign` + foreignKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, foreignPath, registry.SET_VALUE) + require.NoError(t, err, "Should create foreign policy key") + foreignKey.Close() + defer func() { + _ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignPath) + }() + + require.Equal(t, 3, countNRPTRuleKeys(t), "Should have two indexed rules and the legacy one") + + // a configurator that never applied a DNS config, as one built after a + // restart or from a shutdown state without a count is + freshRun := ®istryConfigurator{} + require.NoError(t, freshRun.removeDNSMatchPolicies()) + + assert.Equal(t, 0, countNRPTRuleKeys(t), "Should remove every rule left by the previous run") + + exists, err := registryKeyExists(foreignPath) + require.NoError(t, err) + assert.True(t, exists, "Should not remove a policy that is not ours") +} + +func countNRPTRuleKeys(t *testing.T) int { + t.Helper() + + names, err := listNRPTRuleKeys(DNSPolicyConfigRoot) + require.NoError(t, err, "Should list NRPT rule keys") + return len(names) +} + func cleanupRegistryKeys(*testing.T) { - // Clean up more entries to account for batching tests with many domains - cfg := ®istryConfigurator{nrptEntryCount: 20} + cfg := ®istryConfigurator{} _ = cfg.removeDNSMatchPolicies() } @@ -125,7 +181,7 @@ func TestNRPTDomainBatching(t *testing.T) { // Create a test interface registry key so updateSearchDomains doesn't fail testGUID := "{12345678-1234-1234-1234-123456789ABC}" - interfacePath := `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\` + testGUID + 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() @@ -193,7 +249,7 @@ func TestNRPTDomainBatching(t *testing.T) { require.NoError(t, err) // Verify that exactly expectedRuleCount rules were created - assert.Equal(t, tc.expectedRuleCount, cfg.nrptEntryCount, + assert.Equal(t, tc.expectedRuleCount, countNRPTRuleKeys(t), "Should create %d NRPT rules for %d domains", tc.expectedRuleCount, tc.domainCount) // Verify all expected rules exist diff --git a/client/internal/dns/interface_index.go b/client/internal/dns/interface_index.go new file mode 100644 index 000000000..9e7dca080 --- /dev/null +++ b/client/internal/dns/interface_index.go @@ -0,0 +1,15 @@ +package dns + +import ( + "fmt" + "net" +) + +func getInterfaceIndex(interfaceName string) (int, error) { + iface, err := net.InterfaceByName(interfaceName) + if err != nil { + return 0, fmt.Errorf("lookup interface %q: %w", interfaceName, err) + } + + return iface.Index, nil +} diff --git a/client/internal/dns/interface_index_test.go b/client/internal/dns/interface_index_test.go new file mode 100644 index 000000000..9b146398a --- /dev/null +++ b/client/internal/dns/interface_index_test.go @@ -0,0 +1,35 @@ +package dns + +import ( + "net" + "testing" +) + +func TestGetInterfaceIndexExisting(t *testing.T) { + interfaces, err := net.Interfaces() + if err != nil { + t.Fatalf("list network interfaces: %v", err) + } + if len(interfaces) == 0 { + t.Fatal("expected at least one network interface") + } + + iface := interfaces[0] + index, err := getInterfaceIndex(iface.Name) + if err != nil { + t.Fatalf("look up existing interface %q: %v", iface.Name, err) + } + if index != iface.Index { + t.Fatalf("expected interface index %d, got %d", iface.Index, index) + } +} + +func TestGetInterfaceIndexMissing(t *testing.T) { + index, err := getInterfaceIndex("netbird-interface-that-does-not-exist") + if index != 0 { + t.Fatalf("expected missing interface index to be 0, got %d", index) + } + if err == nil { + t.Fatal("expected missing interface lookup to return an error") + } +} diff --git a/client/internal/dns/local/local.go b/client/internal/dns/local/local.go index d0268186c..fef35fd41 100644 --- a/client/internal/dns/local/local.go +++ b/client/internal/dns/local/local.go @@ -6,6 +6,7 @@ import ( "fmt" "net" "net/netip" + "os" "slices" "strings" "sync" @@ -36,7 +37,43 @@ type resolver interface { // record is left alone (it points at something outside our mesh, e.g. // a non-peer upstream). type PeerConnectivity interface { - IsConnectedByIP(ip string) (known, connected bool) + IsConnectedByIP(ip netip.Addr) (known, connected bool) +} + +// PeerActivator wakes lazy-connection peers on demand. The local resolver calls +// it with the tunnel IPs an answer points at, so a peer that is idle (lazily +// disconnected) starts connecting at DNS-resolution time rather than racing the +// client's first request packet. nil disables warm-up. +type PeerActivator interface { + // ActivatePeersByIP triggers wake-up for the peer(s) owning addrs and blocks + // until one is connected or ctx (a short per-query budget) expires. It is a + // fast no-op for unknown or already-connected addresses. + ActivatePeersByIP(ctx context.Context, addrs []netip.Addr) +} + +const ( + defaultLazyWarmupTimeout = 2 * time.Second + envLazyWarmupTimeout = "NB_DNS_LAZY_WARMUP_TIMEOUT" +) + +// lazyWarmupTimeoutFromEnv returns the per-query budget for waking a +// lazy-connection peer a DNS answer points at. Tunable via +// NB_DNS_LAZY_WARMUP_TIMEOUT (a Go duration). Parsed once at construction time. +func lazyWarmupTimeoutFromEnv() time.Duration { + v := os.Getenv(envLazyWarmupTimeout) + if v == "" { + return defaultLazyWarmupTimeout + } + d, err := time.ParseDuration(v) + if err != nil { + log.Warnf("invalid %s value %q, using default %s: %v", envLazyWarmupTimeout, v, defaultLazyWarmupTimeout, err) + return defaultLazyWarmupTimeout + } + if d <= 0 { + log.Warnf("non-positive %s value %q, using default %s", envLazyWarmupTimeout, v, defaultLazyWarmupTimeout) + return defaultLazyWarmupTimeout + } + return d } type Resolver struct { @@ -51,6 +88,12 @@ type Resolver struct { // filter and preserves the legacy "return whatever is registered" // behaviour for callers that never wire a status source. peerConn PeerConnectivity + // peerActivator, when non-nil, is called at resolution time to warm the + // lazy connection to the peer(s) an answer points at. nil disables warm-up. + peerActivator PeerActivator + // warmupTimeout is the per-query budget for the lazy-connection warm-up + // wait, resolved from the environment once at construction time. + warmupTimeout time.Duration ctx context.Context cancel context.CancelFunc @@ -59,11 +102,12 @@ type Resolver struct { func NewResolver() *Resolver { ctx, cancel := context.WithCancel(context.Background()) return &Resolver{ - records: make(map[dns.Question][]dns.RR), - domains: make(map[domain.Domain]struct{}), - zones: make(map[domain.Domain]bool), - ctx: ctx, - cancel: cancel, + records: make(map[dns.Question][]dns.RR), + domains: make(map[domain.Domain]struct{}), + zones: make(map[domain.Domain]bool), + warmupTimeout: lazyWarmupTimeoutFromEnv(), + ctx: ctx, + cancel: cancel, } } @@ -76,6 +120,14 @@ func (d *Resolver) SetPeerConnectivity(p PeerConnectivity) { d.peerConn = p } +// SetPeerActivator wires the DNS-time lazy-connection warm-up. Pass nil to +// disable. Safe to call multiple times; the latest value wins. +func (d *Resolver) SetPeerActivator(a PeerActivator) { + d.mu.Lock() + defer d.mu.Unlock() + d.peerActivator = a +} + func (d *Resolver) MatchSubdomains() bool { return true } @@ -122,6 +174,9 @@ func (d *Resolver) ServeDNS(w dns.ResponseWriter, r *dns.Msg) { replyMessage.RecursionAvailable = true result := d.lookupRecords(logger, question) + // Warm before filtering: activation flips a lazily-idle target to connected, + // which then lets it survive the disconnected-peer filter below. + d.warmLazyPeers(question, result.records) result.records = d.filterDisconnectedPeerAnswers(logger, question, result.records) replyMessage.Authoritative = !result.hasExternalData replyMessage.Answer = result.records @@ -495,8 +550,8 @@ func (d *Resolver) filterDisconnectedPeerAnswers(logger *log.Entry, question dns kept := make([]dns.RR, 0, len(records)) var dropped int for _, rr := range records { - ip := extractRecordIP(rr) - if ip == "" { + ip, ok := extractRecordAddr(rr) + if !ok { kept = append(kept, rr) continue } @@ -518,22 +573,57 @@ func (d *Resolver) filterDisconnectedPeerAnswers(logger *log.Entry, question dns return kept } -// extractRecordIP returns the dotted-decimal / colon-hex IP carried by -// an A or AAAA record, or "" for any other record type. -func extractRecordIP(rr dns.RR) string { +// warmLazyPeers triggers lazy-connection wake-up for the peers a resolved +// answer points at and waits briefly for one to connect, so the caller's first +// request doesn't race the connection establishment. Warm-up is scoped to +// match-only (non-authoritative) zones — the synthesized private-service zones +// and user-created zones whose records point at specific peers. The account's +// peer zone is authoritative, so plain peer-name lookups never trigger warm-up; +// otherwise resolving any peer's name would wake its idle connection, defeating +// laziness mesh-wide. No-op when no activator is wired (lazy connections +// disabled) or the answer carries no peer IPs. +func (d *Resolver) warmLazyPeers(question dns.Question, records []dns.RR) { + if len(records) < 2 { + return + } + d.mu.RLock() + activator := d.peerActivator + var nonAuth, found bool + if activator != nil { + nonAuth, found = d.findZone(question.Name) + } + d.mu.RUnlock() + if activator == nil || !found || !nonAuth { + return + } + + var addrs []netip.Addr + for _, rr := range records { + if addr, ok := extractRecordAddr(rr); ok { + addrs = append(addrs, addr) + } + } + if len(addrs) == 0 { + return + } + + ctx, cancel := context.WithTimeout(d.ctx, d.warmupTimeout) + defer cancel() + activator.ActivatePeersByIP(ctx, addrs) +} + +// extractRecordAddr returns the IP address carried by an A or AAAA record. +// ok is false for any other record type or a record with no address. +func extractRecordAddr(rr dns.RR) (netip.Addr, bool) { switch r := rr.(type) { case *dns.A: - if r.A == nil { - return "" - } - return r.A.String() + addr, ok := netip.AddrFromSlice(r.A) + return addr.Unmap(), ok case *dns.AAAA: - if r.AAAA == nil { - return "" - } - return r.AAAA.String() + addr, ok := netip.AddrFromSlice(r.AAAA) + return addr.Unmap(), ok } - return "" + return netip.Addr{}, false } // Update replaces all zones and their records diff --git a/client/internal/dns/local/local_test.go b/client/internal/dns/local/local_test.go index 9b7dac231..89e896c0a 100644 --- a/client/internal/dns/local/local_test.go +++ b/client/internal/dns/local/local_test.go @@ -37,8 +37,8 @@ type mockPeerConnectivity struct { byIP map[string]struct{ known, connected bool } } -func (m mockPeerConnectivity) IsConnectedByIP(ip string) (known, connected bool) { - v, ok := m.byIP[ip] +func (m mockPeerConnectivity) IsConnectedByIP(ip netip.Addr) (known, connected bool) { + v, ok := m.byIP[ip.String()] if !ok { return false, false } diff --git a/client/internal/dns/local/warmup_test.go b/client/internal/dns/local/warmup_test.go new file mode 100644 index 000000000..0e77aa963 --- /dev/null +++ b/client/internal/dns/local/warmup_test.go @@ -0,0 +1,204 @@ +package local + +import ( + "context" + "net" + "net/netip" + "sync" + "testing" + "time" + + "github.com/miekg/dns" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/dns/test" + nbdns "github.com/netbirdio/netbird/dns" +) + +// recordingActivator records the addresses it was asked to warm and returns +// immediately, so ServeDNS is not blocked by the test. +type recordingActivator struct { + mu sync.Mutex + called bool + addrs []netip.Addr +} + +func (r *recordingActivator) ActivatePeersByIP(_ context.Context, addrs []netip.Addr) { + r.mu.Lock() + defer r.mu.Unlock() + r.called = true + r.addrs = append(r.addrs, addrs...) +} + +func serveA(t *testing.T, resolver *Resolver, name string) *dns.Msg { + t.Helper() + var resp *dns.Msg + w := &test.MockResponseWriter{WriteMsgFunc: func(m *dns.Msg) error { resp = m; return nil }} + resolver.ServeDNS(w, new(dns.Msg).SetQuestion(name, dns.TypeA)) + return resp +} + +// serviceZone registers rec in a match-only (non-authoritative) zone, the shape +// the synthesized private-service zones arrive in. +func serviceZone(t *testing.T, resolver *Resolver, zone string, records ...nbdns.SimpleRecord) { + t.Helper() + resolver.Update([]nbdns.CustomZone{{ + Domain: zone, + Records: records, + NonAuthoritative: true, + }}) +} + +func TestLocalResolver_WarmsLazyPeerOnResolve(t *testing.T) { + // Warm-up fires only for multi-record answers (the HA / round-robin shape of + // the synthesized private-service zones), so register two peer targets. + const name = "svc.proxy.netbird.cloud." + recs := []nbdns.SimpleRecord{ + {Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"}, + {Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.8"}, + } + resolver := NewResolver() + serviceZone(t, resolver, "proxy.netbird.cloud", recs...) + + act := &recordingActivator{} + resolver.SetPeerActivator(act) + + resp := serveA(t, resolver, name) + require.NotNil(t, resp, "resolver must answer") + require.NotEmpty(t, resp.Answer, "answer must carry the A records") + + act.mu.Lock() + defer act.mu.Unlock() + assert.True(t, act.called, "activator must be invoked for a multi-record service-zone answer") + assert.Contains(t, act.addrs, netip.MustParseAddr("100.64.0.7"), "activator must receive the first peer IP") + assert.Contains(t, act.addrs, netip.MustParseAddr("100.64.0.8"), "activator must receive the second peer IP") +} + +func TestLocalResolver_NoWarmupForSingleRecord(t *testing.T) { + // A single-record answer does not trigger warm-up; the resolver only warms + // multi-record answers. + rec := nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"} + resolver := NewResolver() + serviceZone(t, resolver, "proxy.netbird.cloud", rec) + + act := &recordingActivator{} + resolver.SetPeerActivator(act) + + resp := serveA(t, resolver, rec.Name) + require.NotNil(t, resp, "resolver must answer") + require.NotEmpty(t, resp.Answer, "answer must carry the A record") + + act.mu.Lock() + defer act.mu.Unlock() + assert.False(t, act.called, "activator must not be invoked for a single-record answer") +} + +func TestLocalResolver_NoActivatorNoWarmup(t *testing.T) { + // With no activator wired the resolver behaves exactly as before. + rec := nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"} + resolver := NewResolver() + serviceZone(t, resolver, "proxy.netbird.cloud", rec) + + resp := serveA(t, resolver, rec.Name) + require.NotNil(t, resp, "resolver must still answer without an activator") + require.NotEmpty(t, resp.Answer, "answer must carry the A record") +} + +func TestLocalResolver_NoWarmupForMissingRecord(t *testing.T) { + // A query that resolves to nothing must not invoke the activator (no IPs). + resolver := NewResolver() + serviceZone(t, resolver, "proxy.netbird.cloud", + nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"}) + + act := &recordingActivator{} + resolver.SetPeerActivator(act) + + serveA(t, resolver, "absent.proxy.netbird.cloud.") + + act.mu.Lock() + defer act.mu.Unlock() + assert.False(t, act.called, "activator must not be invoked when there is no answer") +} + +func TestLocalResolver_NoWarmupInAuthoritativeZone(t *testing.T) { + // The account's peer zone is authoritative; resolving a peer's name there + // must not wake its lazy connection — warm-up is scoped to match-only + // (non-authoritative) zones such as the synthesized private-service zones. + // Use a multi-record answer so the authoritative-zone scoping is the only + // reason warm-up is skipped, not the single-record guard. + const name = "peer.netbird.cloud." + recs := []nbdns.SimpleRecord{ + {Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.9"}, + {Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.10"}, + } + resolver := NewResolver() + resolver.Update([]nbdns.CustomZone{{ + Domain: "netbird.cloud", + Records: recs, + }}) + + act := &recordingActivator{} + resolver.SetPeerActivator(act) + + resp := serveA(t, resolver, name) + require.NotNil(t, resp, "resolver must answer") + require.NotEmpty(t, resp.Answer, "answer must carry the A records") + + act.mu.Lock() + defer act.mu.Unlock() + assert.False(t, act.called, "activator must not be invoked for authoritative-zone answers") +} + +func TestLazyWarmupTimeoutFromEnv(t *testing.T) { + tests := []struct { + name string + value string + envSet bool + want time.Duration + }{ + {name: "unset uses default", want: defaultLazyWarmupTimeout}, + {name: "valid overrides", value: "5s", envSet: true, want: 5 * time.Second}, + {name: "invalid falls back", value: "not-a-duration", envSet: true, want: defaultLazyWarmupTimeout}, + {name: "negative falls back", value: "-1s", envSet: true, want: defaultLazyWarmupTimeout}, + {name: "zero falls back", value: "0s", envSet: true, want: defaultLazyWarmupTimeout}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.envSet { + t.Setenv(envLazyWarmupTimeout, tt.value) + } + assert.Equal(t, tt.want, lazyWarmupTimeoutFromEnv()) + assert.Equal(t, tt.want, NewResolver().warmupTimeout, "constructor must resolve the timeout once") + }) + } +} + +func TestExtractRecordAddr(t *testing.T) { + t.Run("A record yields unmapped v4", func(t *testing.T) { + // net.ParseIP returns the 16-byte v4-in-v6 form, the same shape + // miekg/dns stores after parsing an A record; the extracted address + // must compare equal to a plain v4 netip.Addr. + addr, ok := extractRecordAddr(&dns.A{A: net.ParseIP("100.64.0.7")}) + require.True(t, ok) + assert.True(t, addr.Is4()) + assert.Equal(t, netip.MustParseAddr("100.64.0.7"), addr) + }) + + t.Run("AAAA record yields v6", func(t *testing.T) { + addr, ok := extractRecordAddr(&dns.AAAA{AAAA: net.ParseIP("fd00::1")}) + require.True(t, ok) + assert.Equal(t, netip.MustParseAddr("fd00::1"), addr) + }) + + t.Run("A record without address", func(t *testing.T) { + _, ok := extractRecordAddr(&dns.A{}) + assert.False(t, ok) + }) + + t.Run("non-address record", func(t *testing.T) { + _, ok := extractRecordAddr(&dns.CNAME{Target: "target.netbird.cloud."}) + assert.False(t, ok) + }) +} diff --git a/client/internal/dns/mock_server.go b/client/internal/dns/mock_server.go index 31fedd9e5..b19862c2f 100644 --- a/client/internal/dns/mock_server.go +++ b/client/internal/dns/mock_server.go @@ -8,6 +8,7 @@ import ( "github.com/miekg/dns" dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config" + "github.com/netbirdio/netbird/client/internal/dns/local" nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/shared/management/domain" @@ -92,6 +93,11 @@ func (m *MockServer) SetFirewall(Firewall) { // Mock implementation - no-op } +// SetPeerActivator mock implementation of SetPeerActivator from Server interface +func (m *MockServer) SetPeerActivator(local.PeerActivator) { + // Mock implementation - no-op +} + // BeginBatch mock implementation of BeginBatch from Server interface func (m *MockServer) BeginBatch() { // Mock implementation - no-op diff --git a/client/internal/dns/notifier.go b/client/internal/dns/notifier.go index 35cb6ff82..79d924a78 100644 --- a/client/internal/dns/notifier.go +++ b/client/internal/dns/notifier.go @@ -51,7 +51,5 @@ func (n *notifier) notify() { return } - go func(l listener.NetworkChangeListener) { - l.OnNetworkChanged("") - }(n.listener) + n.listener.OnNetworkChanged("") } diff --git a/client/internal/dns/response_writer_test.go b/client/internal/dns/response_writer_test.go index 857964406..bc8416029 100644 --- a/client/internal/dns/response_writer_test.go +++ b/client/internal/dns/response_writer_test.go @@ -4,7 +4,7 @@ import ( "net" "testing" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/google/gopacket" "github.com/google/gopacket/layers" "github.com/miekg/dns" diff --git a/client/internal/dns/server.go b/client/internal/dns/server.go index 7556c66cc..3af912792 100644 --- a/client/internal/dns/server.go +++ b/client/internal/dns/server.go @@ -82,6 +82,7 @@ type Server interface { PopulateManagementDomain(mgmtURL *url.URL) error SetRouteSources(selected, active func() route.HAMap) SetFirewall(Firewall) + SetPeerActivator(local.PeerActivator) } type nsGroupsByDomain struct { @@ -251,7 +252,7 @@ func NewDefaultServerPermanentUpstream( ds.hostsDNSHolder.set(hostsDnsList) ds.permanent = true ds.currentConfig = dnsConfigToHostDNSConfig(config, ds.service.RuntimeIP(), ds.service.RuntimePort()) - ds.searchDomainNotifier = newNotifier(ds.SearchDomains()) + ds.searchDomainNotifier = newNotifier(ds.searchDomains()) ds.searchDomainNotifier.setListener(listener) setServerDns(ds) return ds @@ -491,6 +492,13 @@ func (s *DefaultServer) SetFirewall(fw Firewall) { } } +// SetPeerActivator wires the DNS-time lazy-connection warm-up on the local +// resolver. Injected after the connection manager exists (it does not at +// DNS-server construction time). Pass nil to disable. +func (s *DefaultServer) SetPeerActivator(a local.PeerActivator) { + s.localResolver.SetPeerActivator(a) +} + // Stop stops the server func (s *DefaultServer) Stop() { s.ctxCancel() @@ -594,6 +602,12 @@ func (s *DefaultServer) UpdateDNSServer(serial uint64, update nbdns.Config) erro } func (s *DefaultServer) SearchDomains() []string { + s.mux.Lock() + defer s.mux.Unlock() + return s.searchDomains() +} + +func (s *DefaultServer) searchDomains() []string { var searchDomains []string for _, dConf := range s.currentConfig.Domains { @@ -678,7 +692,7 @@ func (s *DefaultServer) applyConfiguration(update nbdns.Config) error { }() if s.searchDomainNotifier != nil { - s.searchDomainNotifier.onNewSearchDomains(s.SearchDomains()) + s.searchDomainNotifier.onNewSearchDomains(s.searchDomains()) } s.updateNSGroupStates(update.NameServerGroups) @@ -1435,11 +1449,11 @@ type localPeerConnectivity struct { // IsConnectedByIP looks the IP up in the peerstore and surfaces both // the known and connected bits. Used by Resolver.filterDisconnectedPeerAnswers. -func (l localPeerConnectivity) IsConnectedByIP(ip string) (known, connected bool) { +func (l localPeerConnectivity) IsConnectedByIP(ip netip.Addr) (known, connected bool) { if l.status == nil { return false, false } - state, ok := l.status.PeerStateByIP(ip) + state, ok := l.status.PeerStateByIP(ip.String()) if !ok { return false, false } diff --git a/client/internal/dns/server_privileged_test.go b/client/internal/dns/server_privileged_test.go index a03aea169..a17044cf5 100644 --- a/client/internal/dns/server_privileged_test.go +++ b/client/internal/dns/server_privileged_test.go @@ -9,7 +9,7 @@ import ( "os" "testing" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/miekg/dns" "github.com/stretchr/testify/assert" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" diff --git a/client/internal/dns/unclean_shutdown_windows.go b/client/internal/dns/unclean_shutdown_windows.go index 24a9eca50..ab0b2cc63 100644 --- a/client/internal/dns/unclean_shutdown_windows.go +++ b/client/internal/dns/unclean_shutdown_windows.go @@ -5,9 +5,8 @@ import ( ) type ShutdownState struct { - Guid string - GPO bool - NRPTEntryCount int + Guid string + GPO bool } func (s *ShutdownState) Name() string { @@ -16,9 +15,8 @@ func (s *ShutdownState) Name() string { func (s *ShutdownState) Cleanup() error { manager := ®istryConfigurator{ - guid: s.Guid, - gpo: s.GPO, - nrptEntryCount: s.NRPTEntryCount, + guid: s.Guid, + gpo: s.GPO, } if err := manager.restoreUncleanShutdownDNS(); err != nil { diff --git a/client/internal/dns/upstream_ios.go b/client/internal/dns/upstream_ios.go index b989bf0f9..793d87fca 100644 --- a/client/internal/dns/upstream_ios.go +++ b/client/internal/dns/upstream_ios.go @@ -130,8 +130,3 @@ func GetClientPrivate(iface privateClientIface, upstreamIP netip.Addr, dialTimeo } return client, nil } - -func getInterfaceIndex(interfaceName string) (int, error) { - iface, err := net.InterfaceByName(interfaceName) - return iface.Index, err -} diff --git a/client/internal/dns_peer_activator.go b/client/internal/dns_peer_activator.go new file mode 100644 index 000000000..c283d6251 --- /dev/null +++ b/client/internal/dns_peer_activator.go @@ -0,0 +1,76 @@ +package internal + +import ( + "context" + "net/netip" + "time" + + "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/peerstore" +) + +const dnsActivationPollInterval = 50 * time.Millisecond + +// dnsPeerActivator wakes lazy-connection peers from the DNS resolution path. It +// implements dns/local.PeerActivator. DNS queries run on their own goroutines, +// so it only touches state that is safe for concurrent use — ConnMgr.ActivatePeer, +// peerstore.Store and peer.Status — and never takes the engine's syncMsgMux, +// keeping DNS resolution from contending with network-map processing. +type dnsPeerActivator struct { + connMgr *ConnMgr + peerStore *peerstore.Store + status *peer.Status + // ctx is the engine's long-lived context. The connection dial is tied to it + // (not the per-query DNS wait budget) so a handshake that outlasts the wait + // still completes in the background rather than being cancelled at the deadline. + ctx context.Context +} + +// ActivatePeersByIP triggers wake-up for the peer(s) owning addrs and waits +// until one is connected or ctx (the per-query DNS wait budget) expires. +// Activation itself is tied to the engine's long-lived context so the dial +// survives a wait that times out. Unknown or already-connected addresses are +// skipped, so the steady-state (warm) path adds no latency. +func (a *dnsPeerActivator) ActivatePeersByIP(ctx context.Context, addrs []netip.Addr) { + if a == nil || a.connMgr == nil { + return + } + + var pending []string + for _, addr := range addrs { + ip := addr.String() + st, ok := a.status.PeerStateByIP(ip) + if !ok || st.ConnStatus == peer.StatusConnected { + continue + } + conn, ok := a.peerStore.PeerConn(st.PubKey) + if !ok { + continue + } + a.connMgr.ActivatePeer(a.ctx, conn) + pending = append(pending, ip) + } + + if len(pending) == 0 { + return + } + a.waitConnected(ctx, pending) +} + +// waitConnected blocks until any of ips reports a connected peer or ctx expires. +func (a *dnsPeerActivator) waitConnected(ctx context.Context, ips []string) { + ticker := time.NewTicker(dnsActivationPollInterval) + defer ticker.Stop() + for { + for _, ip := range ips { + if st, ok := a.status.PeerStateByIP(ip); ok && st.ConnStatus == peer.StatusConnected { + return + } + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} diff --git a/client/internal/dns_peer_activator_test.go b/client/internal/dns_peer_activator_test.go new file mode 100644 index 000000000..8c3b75e59 --- /dev/null +++ b/client/internal/dns_peer_activator_test.go @@ -0,0 +1,129 @@ +package internal + +import ( + "context" + "net/netip" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/peerstore" +) + +func newTestPeerConn(t *testing.T, key string) *peer.Conn { + t.Helper() + conn, err := peer.NewConn(peer.ConnConfig{ + Key: key, + LocalKey: "local", + WgConfig: peer.WgConfig{ + AllowedIps: []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")}, + }, + }, peer.ServiceDependencies{}) + require.NoError(t, err) + return conn +} + +func newTestDNSPeerActivator(t *testing.T) (*dnsPeerActivator, *peer.Status, *peerstore.Store) { + t.Helper() + status := peer.NewRecorder("https://mgm") + store := peerstore.NewConnStore() + // ConnMgr without Start: the lazy manager is nil, so ActivatePeer is a + // no-op — these tests exercise the activator's skip/wait logic. + connMgr := NewConnMgr(&EngineConfig{}, status, store, nil) + return &dnsPeerActivator{ + connMgr: connMgr, + peerStore: store, + status: status, + ctx: context.Background(), + }, status, store +} + +func TestDNSPeerActivator_NilSafe(t *testing.T) { + var a *dnsPeerActivator + a.ActivatePeersByIP(context.Background(), []netip.Addr{netip.MustParseAddr("100.64.0.1")}) +} + +// TestDNSPeerActivator_SkipsUnknownAndConnectedPeers verifies the steady-state +// (warm) path adds no latency: already-connected and unknown addresses never +// enter the wait loop. +func TestDNSPeerActivator_SkipsUnknownAndConnectedPeers(t *testing.T) { + a, status, store := newTestDNSPeerActivator(t) + + require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", "fd00::1")) + require.NoError(t, status.UpdatePeerState(peer.State{PubKey: "peerA", ConnStatus: peer.StatusConnected})) + store.AddPeerConn("peerA", newTestPeerConn(t, "peerA")) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + start := time.Now() + a.ActivatePeersByIP(ctx, []netip.Addr{ + netip.MustParseAddr("100.64.0.1"), // known, connected -> skipped + netip.MustParseAddr("fd00::1"), // known via IPv6, connected -> skipped + netip.MustParseAddr("100.64.0.99"), // unknown -> skipped + }) + require.Less(t, time.Since(start), time.Second, "no pending peer must mean no wait") +} + +// TestDNSPeerActivator_WaitsForPendingPeerToConnect verifies the wait loop +// returns as soon as a pending peer reports connected, well before the +// per-query budget expires. +func TestDNSPeerActivator_WaitsForPendingPeerToConnect(t *testing.T) { + a, status, store := newTestDNSPeerActivator(t) + + require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", "")) + store.AddPeerConn("peerA", newTestPeerConn(t, "peerA")) + + go func() { + time.Sleep(150 * time.Millisecond) + _ = status.UpdatePeerState(peer.State{PubKey: "peerA", ConnStatus: peer.StatusConnected}) + }() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + start := time.Now() + a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")}) + elapsed := time.Since(start) + + require.GreaterOrEqual(t, elapsed, 100*time.Millisecond, "must wait for the pending peer") + require.Less(t, elapsed, 5*time.Second, "must return on connect, not at the deadline") +} + +// TestDNSPeerActivator_ReturnsAtBudgetWhenPeerStaysIdle verifies a peer that +// never connects releases the DNS response at the per-query budget instead of +// blocking it indefinitely. +func TestDNSPeerActivator_ReturnsAtBudgetWhenPeerStaysIdle(t *testing.T) { + a, status, store := newTestDNSPeerActivator(t) + + require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", "")) + store.AddPeerConn("peerA", newTestPeerConn(t, "peerA")) + + ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond) + defer cancel() + + start := time.Now() + a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")}) + elapsed := time.Since(start) + + require.GreaterOrEqual(t, elapsed, 250*time.Millisecond, "must wait out the budget for a pending peer") + require.Less(t, elapsed, 5*time.Second, "must not block past the budget") +} + +// TestDNSPeerActivator_NoWaitWithoutPeerConn verifies a known-but-idle peer +// with no connection object in the store is not waited on: there is nothing to +// activate, so waiting could only ever time out. +func TestDNSPeerActivator_NoWaitWithoutPeerConn(t *testing.T) { + a, status, _ := newTestDNSPeerActivator(t) + + require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", "")) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + start := time.Now() + a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")}) + require.Less(t, time.Since(start), time.Second, "peer without a conn must not be waited on") +} diff --git a/client/internal/dnsfwd/manager.go b/client/internal/dnsfwd/manager.go index c4c16cd3f..29ca0d247 100644 --- a/client/internal/dnsfwd/manager.go +++ b/client/internal/dnsfwd/manager.go @@ -101,7 +101,7 @@ func (m *Manager) Start(fwdEntries []*ForwarderEntry) error { m.dnsForwarder = NewDNSForwarder(listenAddress, dnsTTL, m.firewall, m.statusRecorder, m.wgIface) go func() { - if err := m.dnsForwarder.Listen(fwdEntries); err != nil { + if err := m.dnsForwarder.Listen(fwdEntries); err != nil { //nolint:staticcheck // todo handle close error if it is exists log.Errorf("failed to start DNS forwarder, err: %v", err) } diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.o b/client/internal/ebpf/ebpf/bpf_bpfeb.o index 6e9cda44a..7433ad740 100644 Binary files a/client/internal/ebpf/ebpf/bpf_bpfeb.o and b/client/internal/ebpf/ebpf/bpf_bpfeb.o differ diff --git a/client/internal/ebpf/ebpf/bpf_bpfel.o b/client/internal/ebpf/ebpf/bpf_bpfel.o index 6338f4774..779f43a00 100644 Binary files a/client/internal/ebpf/ebpf/bpf_bpfel.o and b/client/internal/ebpf/ebpf/bpf_bpfel.o differ diff --git a/client/internal/ebpf/ebpf/src/dns_fwd.c b/client/internal/ebpf/ebpf/src/dns_fwd.c index 5f3fbcc32..9f8de2001 100644 --- a/client/internal/ebpf/ebpf/src/dns_fwd.c +++ b/client/internal/ebpf/ebpf/src/dns_fwd.c @@ -52,11 +52,14 @@ int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) { if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) { udp->dest = dns_port; + // Clear the now-stale checksum; zero means "not computed" for IPv4. + udp->check = 0; return XDP_PASS; } if (udp->source == dns_port && ip->saddr == dns_ip) { udp->source = GENERAL_DNS_PORT; + udp->check = 0; return XDP_PASS; } diff --git a/client/internal/ebpf/ebpf/src/wg_proxy.c b/client/internal/ebpf/ebpf/src/wg_proxy.c index 88fea65cf..5e7474928 100644 --- a/client/internal/ebpf/ebpf/src/wg_proxy.c +++ b/client/internal/ebpf/ebpf/src/wg_proxy.c @@ -50,5 +50,11 @@ int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) { __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/engine.go b/client/internal/engine.go index 1d00ed0d2..fac5224c8 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -23,6 +23,7 @@ import ( "golang.zx2c4.com/wireguard/tun/netstack" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" + "github.com/netbirdio/netbird/client/anonymize" nberrors "github.com/netbirdio/netbird/client/errors" "github.com/netbirdio/netbird/client/firewall" "github.com/netbirdio/netbird/client/firewall/firewalld" @@ -58,13 +59,17 @@ import ( "github.com/netbirdio/netbird/client/internal/syncstore" "github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/jobexec" + "github.com/netbirdio/netbird/client/netevents" cProto "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/system" nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/route" mgm "github.com/netbirdio/netbird/shared/management/client" "github.com/netbirdio/netbird/shared/management/domain" + sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc" + nbnetworkmap "github.com/netbirdio/netbird/shared/management/networkmap" mgmProto "github.com/netbirdio/netbird/shared/management/proto" + types "github.com/netbirdio/netbird/shared/management/types" "github.com/netbirdio/netbird/shared/netiputil" auth "github.com/netbirdio/netbird/shared/relay/auth/hmac" relayClient "github.com/netbirdio/netbird/shared/relay/client" @@ -147,6 +152,7 @@ type EngineConfig struct { BlockLANAccess bool BlockInbound bool DisableIPv6 bool + SyncMessageVersion *int // LazyConnection is the MDM-sourced lazy-connection override; StateUnset defers to // the env var and management feature flag. @@ -176,6 +182,9 @@ type EngineServices struct { UpdateManager *updater.Manager ClientMetrics *metrics.ClientMetrics MetricsCtx context.Context + // NetMgr gates the reconnection loops on OS-reported network + // availability; nil disables gating. + NetMgr *netevents.Manager } // Engine is a mechanism responsible for reacting on Signal and Management stream events and managing connections to the remote peers. @@ -199,6 +208,10 @@ type Engine struct { config *EngineConfig mobileDep MobileDependency + // netMgr gates the peer reconnection guards on OS-reported network + // availability; nil disables gating. + netMgr *netevents.Manager + // STUNs is a list of STUN servers used by ICE STUNs []*stun.URI // TURNs is a list of STUN servers used by ICE @@ -220,6 +233,13 @@ type Engine struct { // networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service networkSerial uint64 + // latestComponents is the most-recent NetworkMapComponents decoded from + // a NetworkMapEnvelope (capability=3 peers only). Held alongside the + // NetworkMap that Calculate() produced from it so future incremental + // updates have a base to apply changes against. nil for legacy-format + // peers. Guarded by syncMsgMux. + latestComponents *types.NetworkMapComponents + networkMonitor *networkmonitor.NetworkMonitor sshServer sshServer @@ -325,6 +345,7 @@ func NewEngine( syncMsgMux: &sync.Mutex{}, config: config, mobileDep: mobileDep, + netMgr: services.NetMgr, STUNs: []*stun.URI{}, TURNs: []*stun.URI{}, networkSerial: 0, @@ -561,12 +582,7 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) } e.stateManager.Start() - initialRoutes, dnsConfig, dnsFeatureFlag, err := e.readInitialSettings() - if err != nil { - return fmt.Errorf("read initial settings: %w", err) - } - - dnsServer, err := e.newDnsServer(dnsConfig) + dnsServer, err := e.newDnsServer() if err != nil { return fmt.Errorf("create dns server: %w", err) } @@ -584,10 +600,8 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) WGInterface: e.wgInterface, StatusRecorder: e.statusRecorder, RelayManager: e.relayManager, - InitialRoutes: initialRoutes, StateManager: e.stateManager, DNSServer: dnsServer, - DNSFeatureFlag: dnsFeatureFlag, PeerStore: e.peerStore, DisableClientRoutes: e.config.DisableClientRoutes, DisableServerRoutes: e.config.DisableServerRoutes, @@ -652,8 +666,24 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) iceCfg := e.createICEConfig() e.connMgr = NewConnMgr(e.config, e.statusRecorder, e.peerStore, wgIface) + e.connMgr.SetRoutedIPsReconciler(func(peerKey string) error { + if e.routeManager == nil { + return nil + } + return e.routeManager.ReconcilePeerAllowedIPs(peerKey) + }) e.connMgr.Start(e.ctx) + // Wire DNS-time lazy-connection warm-up now that the connection manager + // exists (it does not at DNS-server construction time). A DNS answer that + // points at an idle peer then wakes it before the client's first request. + e.dnsServer.SetPeerActivator(&dnsPeerActivator{ + connMgr: e.connMgr, + peerStore: e.peerStore, + status: e.statusRecorder, + ctx: e.ctx, + }) + e.srWatcher = guard.NewSRWatcher(e.signal, e.relayManager, e.mobileDep.IFaceDiscover, iceCfg) e.srWatcher.Start(peer.IsForceRelayed()) @@ -963,8 +993,12 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error { e.ApplySessionDeadline(update.GetSessionExpiresAt()) - if update.NetworkMap != nil && update.NetworkMap.PeerConfig != nil { - e.handleAutoUpdateVersion(update.NetworkMap.PeerConfig.AutoUpdate) + // Envelope sync responses carry PeerConfig at the top level; legacy + // NetworkMap syncs carry it under NetworkMap.PeerConfig. + if pc := update.GetPeerConfig(); pc != nil { + e.handleAutoUpdateVersion(pc.GetAutoUpdate()) + } else if nm := update.GetNetworkMap(); nm != nil && nm.GetPeerConfig() != nil { + e.handleAutoUpdateVersion(nm.GetPeerConfig().GetAutoUpdate()) } done := e.phase("netbird_config") @@ -974,12 +1008,47 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error { return err } + // Decode the network map from either the components envelope or the + // legacy proto.NetworkMap before the posture-check gating below, so the + // "is there a network map" decision covers both wire shapes. + var ( + nm *mgmProto.NetworkMap + components *types.NetworkMapComponents + ) + if version := update.GetVersion(); version == int32(sharedgrpc.ComponentNetworkMap) { + // Components-format peer: decode the envelope back to typed + // components, run Calculate() locally, and convert to the wire + // NetworkMap shape the rest of the engine consumes. Components are + // retained so future incremental updates can apply deltas instead + // of doing a full reconstruction. + envelope := update.GetNetworkMapEnvelope() + if envelope == nil { + return fmt.Errorf("received a SyncReponse indicating use of components network map, but components are missing") + } + + localKey := e.config.WgPrivateKey.PublicKey().String() + dnsName := "" + if pc := update.GetPeerConfig(); pc != nil { + // PeerConfig.Fqdn = "." — extract the + // shared domain by stripping the peer's own label prefix. Falls + // 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) + if err != nil { + return fmt.Errorf("decode network map envelope: %w", err) + } + nm = result.NetworkMap + components = result.Components + } else { + nm = update.GetNetworkMap() + } + // Posture checks are bound to the network map presence: // NetworkMap != nil, checks present -> apply the received checks // NetworkMap != nil, checks nil -> posture checks were removed, clear them // NetworkMap == nil -> config-only update (e.g. relay token rotation), // leave the previously applied checks untouched - nm := update.GetNetworkMap() if nm == nil { return nil } @@ -992,6 +1061,14 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error { } done = e.phase("persist") + // Only retain the components view when the server sent the envelope + // path. A legacy proto.NetworkMap means components == nil; writing it + // here would clobber a previously-cached snapshot, breaking the + // incremental-delta base on a future envelope sync. + if components != nil { + e.latestComponents = components + } + e.persistSyncResponse(update) done() @@ -1005,6 +1082,19 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error { return nil } +// extractDNSDomainFromFQDN returns the trailing dotted domain part of the +// receiving peer's FQDN — the same value the management server fills as +// dnsName when it builds the legacy NetworkMap. "peer42.netbird.cloud" → +// "netbird.cloud". An empty string is returned for unrecognized formats. +func extractDNSDomainFromFQDN(fqdn string) string { + for i := 0; i < len(fqdn); i++ { + if fqdn[i] == '.' && i+1 < len(fqdn) { + return fqdn[i+1:] + } + } + return "" +} + // updateNetbirdConfig applies the management-provided NetBird configuration: // STUN/TURN and relay servers, flow logging and DNS settings. A nil config is a no-op, // which is the case for sync updates carrying only a network map. @@ -1164,6 +1254,7 @@ func (e *Engine) applyInfoFlags(info *system.Info) { e.config.BlockLANAccess, e.config.BlockInbound, e.config.DisableIPv6, + e.config.SyncMessageVersion, e.config.EnableSSHRoot, e.config.EnableSSHSFTP, e.config.EnableSSHLocalPortForwarding, @@ -1304,6 +1395,7 @@ func (e *Engine) handleBundle(params *mgmProto.BundleParameters) (*mgmProto.JobR bundleJobParams := debug.BundleConfig{ Anonymize: params.Anonymize, + AnonymizeLevel: anonymize.ParseLevel(params.AnonymizeLevel), IncludeSystemInfo: true, LogFileCount: uint32(params.LogFileCount), } @@ -1811,6 +1903,7 @@ func (e *Engine) createPeerConn(pubKey string, allowedIPs []netip.Prefix, agentV PermissiveMode: e.config.RosenpassPermissive, }, ICEConfig: e.createICEConfig(), + NetMgr: e.netMgr, } serviceDependencies := peer.ServiceDependencies{ @@ -2014,41 +2107,6 @@ func (e *Engine) close() { } } -func (e *Engine) readInitialSettings() ([]*route.Route, *nbdns.Config, bool, error) { - if runtime.GOOS != "android" { - // nolint:nilnil - return nil, nil, false, nil - } - - info := system.GetInfo(e.ctx) - info.SetFlags( - e.config.RosenpassEnabled, - e.config.RosenpassPermissive, - &e.config.ServerSSHAllowed, - e.config.DisableClientRoutes, - e.config.DisableServerRoutes, - e.config.DisableDNS, - e.config.DisableFirewall, - e.config.BlockLANAccess, - e.config.BlockInbound, - e.config.DisableIPv6, - e.config.EnableSSHRoot, - e.config.EnableSSHSFTP, - e.config.EnableSSHLocalPortForwarding, - e.config.EnableSSHRemotePortForwarding, - e.config.DisableSSHAuth, - ) - - netMap, err := e.mgmClient.GetNetworkMap(info) - if err != nil { - return nil, nil, false, err - } - routes := toRoutes(netMap.GetRoutes()) - dnsCfg := toDNSConfig(netMap.GetDNSConfig(), e.wgInterface.Address()) - dnsFeatureFlag := toDNSFeatureFlag(netMap) - return routes, &dnsCfg, dnsFeatureFlag, nil -} - func (e *Engine) newWgIface() (*iface.WGIface, error) { transportNet, err := e.newStdNet() if err != nil { @@ -2083,7 +2141,7 @@ func (e *Engine) newWgIface() (*iface.WGIface, error) { func (e *Engine) wgInterfaceCreate() (err error) { switch runtime.GOOS { case "android": - err = e.wgInterface.CreateOnAndroid(e.routeManager.InitialRouteRange(), e.dnsServer.DnsIP().String(), e.dnsServer.SearchDomains()) + err = e.wgInterface.CreateOnAndroid(e.routeManager.CurrentRouteRange(), e.dnsServer.DnsIP().String(), e.dnsServer.SearchDomains()) case "ios": e.mobileDep.NetworkChangeListener.SetInterfaceIP(e.config.WgAddr.String()) if e.config.WgAddr.HasIPv6() { @@ -2096,7 +2154,7 @@ func (e *Engine) wgInterfaceCreate() (err error) { return err } -func (e *Engine) newDnsServer(dnsConfig *nbdns.Config) (dns.Server, error) { +func (e *Engine) newDnsServer() (dns.Server, error) { // due to tests where we are using a mocked version of the DNS server if e.dnsServer != nil { return e.dnsServer, nil @@ -2108,7 +2166,7 @@ func (e *Engine) newDnsServer(dnsConfig *nbdns.Config) (dns.Server, error) { e.ctx, e.wgInterface, e.mobileDep.HostDNSAddresses, - *dnsConfig, + nbdns.Config{}, e.mobileDep.NetworkChangeListener, e.statusRecorder, e.config.DisableDNS, @@ -2514,7 +2572,7 @@ func (e *Engine) SetCapture(pc device.PacketCapture) error { } afc := capture.NewAFPacketCapture(intf.Name(), sess) - if err := afc.Start(); err != nil { + if err := afc.Start(); err != nil { //nolint:staticcheck // always errors on non-Linux builds return fmt.Errorf("start AF_PACKET capture on %s: %w", intf.Name(), err) } e.afpacketCapture = afc diff --git a/client/internal/engine_privileged_test.go b/client/internal/engine_privileged_test.go index f787f741f..032992464 100644 --- a/client/internal/engine_privileged_test.go +++ b/client/internal/engine_privileged_test.go @@ -12,7 +12,7 @@ import ( "testing" "time" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/google/uuid" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" diff --git a/client/internal/engine_session_deadline_test.go b/client/internal/engine_session_deadline_test.go index 6127e5bb0..5a67f103a 100644 --- a/client/internal/engine_session_deadline_test.go +++ b/client/internal/engine_session_deadline_test.go @@ -75,4 +75,14 @@ func TestApplySessionDeadline_ThreeState(t *testing.T) { require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(), "invalid timestamp must clear the deadline") }) + + t.Run("recently expired timestamp stays visible as expired", func(t *testing.T) { + e := newEngine() + expired := time.Now().Add(-5 * time.Minute).UTC().Truncate(time.Second) + + e.ApplySessionDeadline(timestamppb.New(expired)) + + require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(expired), + "recently-expired deadline must stay on the recorder so consumers render it as expired") + }) } diff --git a/client/internal/engine_tunsettings.go b/client/internal/engine_tunsettings.go new file mode 100644 index 000000000..34a59671a --- /dev/null +++ b/client/internal/engine_tunsettings.go @@ -0,0 +1,20 @@ +package internal + +func (e *Engine) TunSettings() ([]string, []string) { + e.syncMsgMux.Lock() + routeManager := e.routeManager + dnsServer := e.dnsServer + e.syncMsgMux.Unlock() + + var routes []string + if routeManager != nil { + routes = routeManager.CurrentRouteRange() + } + + var searchDomains []string + if dnsServer != nil { + searchDomains = dnsServer.SearchDomains() + } + + return routes, searchDomains +} diff --git a/client/internal/ipcauth/creds_stub.go b/client/internal/ipcauth/creds_stub.go new file mode 100644 index 000000000..154948716 --- /dev/null +++ b/client/internal/ipcauth/creds_stub.go @@ -0,0 +1,31 @@ +//go:build !linux && !darwin && !freebsd && !windows + +package ipcauth + +import ( + "errors" + "net" + + "google.golang.org/grpc/credentials" +) + +// errUnsupported is returned on platforms with no local peer-identity +// primitive, so consumers fail closed instead of guessing an identity. +var errUnsupported = errors.New("peer identity is not available on this platform") + +// NewTransportCredentials returns nil: without a peer-identity primitive the +// daemon cannot authenticate local callers, and the caller must treat that as +// "authorization cannot be enforced". +func NewTransportCredentials() credentials.TransportCredentials { + return nil +} + +// PeerIdentity always fails on this platform. +func PeerIdentity(net.Conn) (Identity, error) { + return Identity{}, errUnsupported +} + +// ConnIdentity always fails on this platform. +func ConnIdentity(net.Conn) (Identity, error) { + return Identity{}, errUnsupported +} diff --git a/client/internal/ipcauth/creds_unix.go b/client/internal/ipcauth/creds_unix.go new file mode 100644 index 000000000..688fe4623 --- /dev/null +++ b/client/internal/ipcauth/creds_unix.go @@ -0,0 +1,56 @@ +//go:build linux || darwin || freebsd + +package ipcauth + +import ( + "context" + "net" + + "google.golang.org/grpc/credentials" +) + +// NewTransportCredentials returns gRPC transport credentials that expose the +// caller's kernel-authenticated identity via IdentityFromContext. It returns +// nil on platforms that have no peer-identity primitive, which the caller must +// treat as "authorization cannot be enforced". +// +// The handshake exchanges no bytes on the wire, so a client dialing with +// insecure credentials interoperates with a server using these. That keeps +// older CLI and UI binaries working against an upgraded daemon. +func NewTransportCredentials() credentials.TransportCredentials { + return unixCreds{} +} + +// ConnIdentity extracts the caller's identity from an accepted local IPC +// connection. It is shared by the gRPC transport credentials and by the JSON +// gateway, which reads the identity of its own HTTP clients. +func ConnIdentity(conn net.Conn) (Identity, error) { + return PeerIdentity(conn) +} + +type unixCreds struct{} + +func (unixCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + return conn, AuthInfo{}, nil +} + +// ServerHandshake extracts the peer identity and fails closed when it cannot +// be read, so a connection whose caller is unknown never reaches a handler. +func (unixCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + id, err := ConnIdentity(conn) + if err != nil { + return nil, nil, err + } + return conn, AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, nil +} + +func (unixCreds) Info() credentials.ProtocolInfo { + return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()} +} + +func (unixCreds) Clone() credentials.TransportCredentials { return unixCreds{} } + +func (unixCreds) OverrideServerName(string) error { return nil } diff --git a/client/internal/ipcauth/creds_windows.go b/client/internal/ipcauth/creds_windows.go new file mode 100644 index 000000000..37f902c52 --- /dev/null +++ b/client/internal/ipcauth/creds_windows.go @@ -0,0 +1,194 @@ +//go:build windows + +package ipcauth + +import ( + "context" + "fmt" + "net" + "runtime" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" + "google.golang.org/grpc/credentials" +) + +var ( + modadvapi32 = windows.NewLazySystemDLL("advapi32.dll") + procImpersonateNamedPipeClient = modadvapi32.NewProc("ImpersonateNamedPipeClient") +) + +// DefaultPipeSDDL is the security descriptor for the daemon control pipe. +// +// D:P protected DACL, no inheritance +// (A;;GA;;;SY) allow GENERIC_ALL to LocalSystem (the daemon's service account) +// (A;;GA;;;WD) allow GENERIC_ALL to Everyone +// +// Any local caller may connect, as with a Unix socket at 0666; what a caller may +// actually do is decided from its token, not from the DACL. Remote callers are not +// a concern here: winio.ListenPipe creates the pipe with +// FILE_PIPE_REJECT_REMOTE_CLIENTS, so NPFS rejects connections from other machines +// before the descriptor is consulted. +// +// A deny ACE on the NETWORK SID would not add anything and would break callers: +// that SID is present in any network-logon token, which includes OpenSSH and WinRM +// sessions, so it denies administrators driving the CLI over SSH and denies the +// daemon itself when started from such a session. +func DefaultPipeSDDL() string { + return "D:P(A;;GA;;;SY)(A;;GA;;;WD)" +} + +// NewTransportCredentials returns gRPC transport credentials that derive the +// caller's identity from the named-pipe client token. +// +// The client must connect at SECURITY_IDENTIFICATION for the daemon to be able +// to read its token, which is what DialNamedPipe does. +func NewTransportCredentials() credentials.TransportCredentials { + return winpipeCreds{} +} + +// ConnIdentity extracts the caller's identity from an accepted named-pipe +// connection by impersonating the pipe client and reading its token. It is +// shared by the gRPC transport credentials and by the JSON gateway, which +// reads the identity of its own HTTP clients. +func ConnIdentity(conn net.Conn) (Identity, error) { + // go-winio's pipe connection embeds *win32File, which exposes Fd(). + fdConn, ok := conn.(interface{ Fd() uintptr }) + if !ok { + return Identity{}, fmt.Errorf("connection %T does not expose a pipe handle", conn) + } + return pipeClientIdentity(windows.Handle(fdConn.Fd())) +} + +type winpipeCreds struct{} + +func (winpipeCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + return conn, AuthInfo{}, nil +} + +// ServerHandshake extracts the connecting client's identity and fails closed +// when the handle or token cannot be read, so a connection whose caller is +// unknown never reaches a handler. +func (winpipeCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + id, err := ConnIdentity(conn) + if err != nil { + return nil, nil, err + } + return conn, AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, nil +} + +func (winpipeCreds) Info() credentials.ProtocolInfo { + return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()} +} + +func (winpipeCreds) Clone() credentials.TransportCredentials { return winpipeCreds{} } + +func (winpipeCreds) OverrideServerName(string) error { return nil } + +// pipeClientIdentity reads the connecting client's user SID, usable group +// SIDs, and elevation state by impersonating the pipe client on this thread +// and reading the resulting impersonation token. +func pipeClientIdentity(handle windows.Handle) (id Identity, err error) { + // Impersonation is per-thread, so the goroutine must stay on this thread + // until RevertToSelf, otherwise an unrelated goroutine could inherit the + // impersonated context. + runtime.LockOSThread() + + // The thread only goes back to the runtime's pool once it is provably no + // longer impersonating the client. If the revert fails, leaving it locked + // makes Go terminate it when this goroutine exits, which costs one thread + // and keeps a thread running as the client from ever being reused. + clean := false + defer func() { + if clean { + runtime.UnlockOSThread() + } + }() + + if err = impersonateNamedPipeClient(handle); err != nil { + clean = true + return Identity{}, fmt.Errorf("impersonate named pipe client: %w", err) + } + defer func() { + // Surface the revert failure only when nothing else failed: leaving + // the thread impersonated is worse than the original error. + revErr := windows.RevertToSelf() + if revErr != nil { + if err == nil { + err = fmt.Errorf("revert impersonation: %w", revErr) + } + return + } + clean = true + }() + + // openAsSelf=true opens the token with the daemon's own process context + // rather than the impersonated client's, so the open cannot fail because + // the client lacks access to its own token. + var token windows.Token + if err = windows.OpenThreadToken(windows.CurrentThread(), windows.TOKEN_QUERY, true, &token); err != nil { + return Identity{}, fmt.Errorf("open thread token: %w", err) + } + defer func() { + if cerr := token.Close(); cerr != nil { + log.Debugf("close client token: %v", cerr) + } + }() + + return identityFromToken(token) +} + +// identityFromToken reads the user SID, usable group SIDs and elevation state +// out of a Windows token. +func identityFromToken(token windows.Token) (Identity, error) { + user, err := token.GetTokenUser() + if err != nil { + return Identity{}, fmt.Errorf("read token user: %w", err) + } + + groups, err := tokenGroupSIDs(token) + if err != nil { + return Identity{}, err + } + + return Identity{ + SID: user.User.Sid.String(), + Groups: groups, + Elevated: token.IsElevated(), + }, nil +} + +// tokenGroupSIDs returns the SIDs of the groups the token can actually +// exercise. Groups that are disabled or marked deny-only are skipped: a +// UAC-filtered administrator carries BUILTIN\Administrators as deny-only, and +// treating that as membership would hand every admin account privilege it +// cannot currently use. +func tokenGroupSIDs(token windows.Token) ([]string, error) { + tg, err := token.GetTokenGroups() + if err != nil { + return nil, fmt.Errorf("read token groups: %w", err) + } + + var sids []string + for _, g := range tg.AllGroups() { + if g.Attributes&windows.SE_GROUP_ENABLED == 0 { + continue + } + if g.Attributes&windows.SE_GROUP_USE_FOR_DENY_ONLY != 0 { + continue + } + sids = append(sids, g.Sid.String()) + } + return sids, nil +} + +func impersonateNamedPipeClient(h windows.Handle) error { + r, _, e := procImpersonateNamedPipeClient.Call(uintptr(h)) + if r == 0 { + return e + } + return nil +} diff --git a/client/internal/ipcauth/forward.go b/client/internal/ipcauth/forward.go new file mode 100644 index 000000000..57749528c --- /dev/null +++ b/client/internal/ipcauth/forward.go @@ -0,0 +1,272 @@ +package ipcauth + +import ( + "context" + "crypto/rand" + "crypto/subtle" + "encoding/hex" + "fmt" + "slices" + "strconv" + "strings" + + "google.golang.org/grpc/metadata" +) + +// Metadata keys the local JSON gateway uses to forward the identity of its own +// HTTP client to the daemon. The gateway runs inside the daemon process and +// re-dials the daemon over the control socket, so without forwarding every +// JSON request would appear to come from the daemon itself. +const ( + // mdFwd marks a request as forwarded by the JSON gateway. It is always + // set, even when the gateway could not read its client's identity, so the + // daemon can tell "no identity available" apart from "not forwarded". + mdFwd = "x-netbird-fwd" + mdFwdUID = "x-netbird-fwd-uid" // Unix user ID + mdFwdGID = "x-netbird-fwd-gid" // Unix primary group ID + mdFwdSID = "x-netbird-fwd-sid" // Windows user SID + mdFwdGroup = "x-netbird-fwd-group" // Windows group SID, repeated + mdFwdElevated = "x-netbird-fwd-elevated" // Windows, "1" when elevated + + // mdFwdProof proves the forwarded identity was stamped by this process. The + // gateway runs inside the daemon, so a secret held in memory is available to + // the only legitimate producer and to nothing else. + mdFwdProof = "x-netbird-fwd-proof" +) + +// forwardKeys is every metadata key the gateway sets. An HTTP client must never +// be able to supply one itself: see IsReservedForwardKey. +var forwardKeys = []string{mdFwd, mdFwdUID, mdFwdGID, mdFwdSID, mdFwdGroup, mdFwdElevated, mdFwdProof} + +// forwardProof authenticates the gateway's forwarding metadata. It is generated +// once per daemon process and never leaves it: it is not written to disk, not +// logged, and not sent anywhere except over the daemon's own control socket to +// itself. +// +// Without it, trusting a forwarded identity rests on every layer in front of it +// stripping incoming forwarding keys, and on each key's value shape being +// distinguishable from an injected one. A single injected group SID or an +// injected "elevated" flag has the same shape as a legitimate one, so no +// cardinality rule can catch it. Requiring the proof means metadata that did not +// come from this process is refused whatever it contains. +var forwardProof = mustForwardProof() + +func mustForwardProof() string { + var buf [32]byte + if _, err := rand.Read(buf[:]); err != nil { + // Continuing would leave the forwarded path authenticated by a + // predictable value, which is worse than not starting. + panic(fmt.Sprintf("generate identity forwarding proof: %v", err)) + } + return hex.EncodeToString(buf[:]) +} + +// IsReservedForwardKey reports whether a gRPC metadata key belongs to the +// gateway's identity forwarding, and therefore must be dropped when it arrives +// from outside. +// +// grpc-gateway maps "Grpc-Metadata-" request headers into gRPC metadata and +// joins them ahead of the values its own annotators add. Without dropping these, +// an HTTP client could hand the daemon "x-netbird-fwd-uid: 0" and be believed, +// because the daemon trusts forwarded metadata when the transport peer is the +// (privileged) gateway. +func IsReservedForwardKey(key string) bool { + key = strings.ToLower(key) + return slices.Contains(forwardKeys, key) +} + +// ForwardIdentityMetadata encodes an HTTP client's identity for the JSON +// gateway to forward to the daemon. When known is false only the marker is +// set, which makes the daemon treat the caller as unidentified rather than as +// the daemon itself. +func ForwardIdentityMetadata(id Identity, known bool) metadata.MD { + md := metadata.MD{} + md.Set(mdFwd, "1") + md.Set(mdFwdProof, forwardProof) + if !known { + return md + } + + if id.IsWindows() { + md.Set(mdFwdSID, id.SID) + if len(id.Groups) > 0 { + md.Set(mdFwdGroup, id.Groups...) + } + if id.Elevated { + md.Set(mdFwdElevated, "1") + } + return md + } + + md.Set(mdFwdUID, strconv.FormatUint(uint64(id.UID), 10)) + md.Set(mdFwdGID, strconv.FormatUint(uint64(id.GID), 10)) + return md +} + +// CallerIdentity returns the identity to authorize a request against. For a +// direct connection that is the transport peer's kernel identity. For a +// request relayed by the local JSON gateway it is the identity the gateway +// forwarded, since the transport peer is then the daemon itself. +// +// A forwarded identity is only honoured when the transport peer is the daemon's +// own identity and the metadata carries this process's forwarding proof, so +// forged forwarding metadata gains a caller nothing. A forwarded request that +// carries no identity is reported as unidentified, never as the daemon. +// +// The second return value is false when no identity could be established, and +// callers MUST fail closed in that case. +func CallerIdentity(ctx context.Context) (Identity, bool) { + id, ok := IdentityFromContext(ctx) + if !ok { + return Identity{}, false + } + + // A forwarding key that arrives more than once did not come from the gateway + // alone, so nothing about the request can be trusted to describe its caller. + // Refusing outright matters because the alternative reading, "not forwarded", + // would authorize the request as the transport peer, which on the gateway's + // connection is the daemon itself. + if duplicatedForwardKey(ctx) { + return Identity{}, false + } + + forwarded := isForwarded(ctx) + + // Our own process on the other end of the socket is the JSON gateway, the only + // thing that dials the daemon from inside it. Such a call must carry a + // forwarded identity; without one there is no caller to authorize, and + // treating it as the daemon would authorize whatever reached the JSON socket. + // Only Linux reports the peer PID, so this is a belt on top of the gateway's + // interceptor rather than the sole guarantee. + if id.PID != 0 && int(id.PID) == selfPID && !forwarded { + return Identity{}, false + } + + // Only the gateway's own connection may speak for someone else. Being + // privileged is not enough and not the point: the gateway runs inside the + // daemon, so it dials as the daemon's identity whatever user that is, which + // also covers a rootless container. + if !forwarded || !IsDaemonSelf(id) { + return id, true + } + + // Speaking for someone else additionally requires the proof only this process + // holds. Refusing is the only safe reading: the transport peer here is the + // daemon itself, so falling back to it would authorize the request as the + // daemon. This is also what makes the forwarded values trustworthy once + // accepted, so they need no shape checks of their own. + if !authenticForward(ctx) { + return Identity{}, false + } + + return forwardedIdentity(ctx) +} + +// duplicatedForwardKey reports whether any forwarding key carries more than one +// value. The gateway's interceptor sets each key exactly once and replaces what +// was already there, so a repeat means a second source supplied it. +func duplicatedForwardKey(ctx context.Context) bool { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return false + } + for _, key := range forwardKeys { + // Group SIDs are legitimately repeated; the rest identify the caller. + if key == mdFwdGroup { + continue + } + if len(md.Get(key)) > 1 { + return true + } + } + return false +} + +// authenticForward reports whether the request carries this process's forwarding +// proof, which only the in-process JSON gateway can supply. +func authenticForward(ctx context.Context) bool { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return false + } + got := mdSingle(md, mdFwdProof) + return subtle.ConstantTimeCompare([]byte(got), []byte(forwardProof)) == 1 +} + +// isForwarded reports whether the request carries the JSON gateway marker. +func isForwarded(ctx context.Context) bool { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return false + } + return mdSingle(md, mdFwd) != "" +} + +// forwardedIdentity decodes the identity the JSON gateway attached. +func forwardedIdentity(ctx context.Context) (Identity, bool) { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return Identity{}, false + } + + if sid := mdSingle(md, mdFwdSID); sid != "" { + return Identity{ + SID: sid, + // Repeated by design, one value per group, and only reachable once + // the forwarding proof has been verified. + Groups: md.Get(mdFwdGroup), + Elevated: mdSingle(md, mdFwdElevated) == "1", + }, true + } + + uid, err := strconv.ParseUint(mdSingle(md, mdFwdUID), 10, 32) + if err != nil { + return Identity{}, false + } + + id := Identity{UID: uint32(uid)} + if gid, err := strconv.ParseUint(mdSingle(md, mdFwdGID), 10, 32); err == nil { + id.GID = uint32(gid) + } + return id, true +} + +// mdSingle returns the value of a forwarded key only when exactly one was +// supplied. The gateway's interceptor sets each key exactly once, so more than one +// value means something else also supplied it, and the whole identity is treated as +// unknown rather than picking a winner. Defence in depth behind the gateway's +// header filter. +func mdSingle(md metadata.MD, key string) string { + if v := md.Get(key); len(v) == 1 { + return v[0] + } + return "" +} + +// WithForwardedIdentity stamps id onto a context's outgoing metadata for the JSON +// gateway's call to the daemon, replacing any forwarding keys already present so +// values supplied from outside cannot survive alongside it. +// +// This is deliberately not done with runtime.WithMetadata: grpc-gateway skips its +// annotators entirely when no request header maps to metadata ("if len(pairs) == 0 +// { return ctx, nil, nil }", runtime/context.go), which an HTTP/1.0 request with no +// Host header over a unix socket achieves. The daemon would then see an unmarked +// call whose transport peer is the daemon's own identity, and authorize it as the +// daemon. A client interceptor runs for every RPC regardless of headers. +func WithForwardedIdentity(ctx context.Context, id Identity, known bool) context.Context { + md, ok := metadata.FromOutgoingContext(ctx) + if !ok { + md = metadata.MD{} + } else { + md = md.Copy() + } + + for _, key := range forwardKeys { + delete(md, key) + } + for key, values := range ForwardIdentityMetadata(id, known) { + md[key] = values + } + + return metadata.NewOutgoingContext(ctx, md) +} diff --git a/client/internal/ipcauth/forward_test.go b/client/internal/ipcauth/forward_test.go new file mode 100644 index 000000000..d9adf05da --- /dev/null +++ b/client/internal/ipcauth/forward_test.go @@ -0,0 +1,214 @@ +package ipcauth + +import ( + "context" + "testing" + + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/peer" +) + +// transportCtx builds a request context as the daemon's transport credentials +// would: the identity of whoever opened the socket, plus whatever metadata the +// request carried. +func transportCtx(id Identity, md metadata.MD) context.Context { + ctx := peer.NewContext(context.Background(), &peer.Peer{ + AuthInfo: AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, + }) + if md != nil { + ctx = metadata.NewIncomingContext(ctx, md) + } + return ctx +} + +var ( + root = Identity{UID: 0} + unprivUser = Identity{UID: 1000, GID: 1000} +) + +// asDaemon pins which identity counts as this process for the duration of a test. +// Without it the test binary's own uid decides, which silently changes what +// "the gateway" means. +func asDaemon(t *testing.T, id Identity) { + t.Helper() + prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate + t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) + selfIdentity, selfKnown = id, true + selfMayDelegate = !id.IsPrivileged() +} + +func TestCallerIdentity_DirectConnections(t *testing.T) { + t.Run("no transport credentials is not an identity", func(t *testing.T) { + if _, ok := CallerIdentity(context.Background()); ok { + t.Fatal("a caller with no credentials must not be identified") + } + }) + + t.Run("a direct caller is its transport identity", func(t *testing.T) { + id, ok := CallerIdentity(transportCtx(unprivUser, nil)) + if !ok || id.UID != 1000 { + t.Fatalf("got %v ok=%t, want uid 1000", id, ok) + } + }) + + // The whole point of honouring forwarded metadata only from a privileged + // transport peer: an unprivileged caller can set any metadata it likes on its + // own connection to the daemon socket. + t.Run("an unprivileged caller cannot forge an identity", func(t *testing.T) { + asDaemon(t, root) + forged := metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdGID, "0") + id, ok := CallerIdentity(transportCtx(unprivUser, forged)) + if !ok { + t.Fatal("caller should still be identified, as itself") + } + if id.IsPrivileged() || id.UID != 1000 { + t.Fatalf("forged metadata was believed: got %v", id) + } + }) +} + +func TestCallerIdentity_GatewayForwarding(t *testing.T) { + t.Run("the gateway's client identity is used, not the gateway's own", func(t *testing.T) { + asDaemon(t, root) + md := ForwardIdentityMetadata(unprivUser, true) + id, ok := CallerIdentity(transportCtx(root, md)) + if !ok { + t.Fatal("forwarded identity should be usable") + } + if id.IsPrivileged() || id.UID != 1000 { + t.Fatalf("got %v, want the forwarded uid 1000 and not privileged", id) + } + }) + + t.Run("a privileged gateway client stays privileged", func(t *testing.T) { + asDaemon(t, root) + md := ForwardIdentityMetadata(root, true) + id, ok := CallerIdentity(transportCtx(root, md)) + if !ok || !id.IsPrivileged() { + t.Fatalf("got %v ok=%t, want a privileged identity", id, ok) + } + }) + + // A JSON socket the gateway cannot read peer credentials from (a TCP socket, + // say) must not make every request look like the daemon itself. + t.Run("an unreadable client identity is unknown, not the daemon", func(t *testing.T) { + asDaemon(t, root) + md := ForwardIdentityMetadata(Identity{}, false) + if _, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatal("a forwarded request with no identity must not be identified") + } + }) + + // grpc-gateway turns Grpc-Metadata- headers into gRPC metadata and joins + // them ahead of its annotators' values. If an HTTP client's header survived + // that, this is the shape the daemon would see: the attacker's uid 0 first, + // the real uid second. The gateway filters those headers out, and reading a + // duplicated key as unknown makes the daemon safe even if it did not. + t.Run("a duplicated key from an injected header is not believed", func(t *testing.T) { + asDaemon(t, root) + md := metadata.MD{} + md.Append(mdFwd, "1") + md.Append(mdFwdUID, "0") // injected by the HTTP client + md.Append(mdFwdUID, "1000") // appended by the gateway's annotator + if id, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatalf("injected uid was accepted: got %v", id) + } + }) + + t.Run("a duplicated marker is not believed either", func(t *testing.T) { + asDaemon(t, root) + md := metadata.MD{} + md.Append(mdFwd, "1") + md.Append(mdFwd, "1") + md.Append(mdFwdUID, "1000") + // A repeated marker must not be read as "not forwarded": that would + // authorize the request as the transport peer, which on the gateway's + // connection is the daemon itself. + if id, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatalf("a duplicated marker was believed: got %v", id) + } + }) + + // The layers in front of this (the gateway's header matcher, and its + // interceptor replacing every forwarding key) are what keep outside metadata + // from arriving at all. The proof is what the daemon can check for itself, and + // it is the only defence that works for a value whose legitimate shape is + // indistinguishable from an injected one: a lone group SID, or "elevated". + t.Run("forwarding metadata without this process's proof is refused", func(t *testing.T) { + asDaemon(t, root) + for name, md := range map[string]metadata.MD{ + "no proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0"), + "wrong proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdProof, "deadbeef"), + "windows identity without a proof": metadata.Pairs(mdFwd, "1", + mdFwdSID, "S-1-5-21-1-2-3-1001", mdFwdGroup, sidAdministrators, mdFwdElevated, "1"), + } { + t.Run(name, func(t *testing.T) { + if id, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatalf("unstamped forwarding metadata was believed: got %v", id) + } + }) + } + }) + + // A caller that reaches the gateway cannot see the proof, so it cannot append + // a group of its own to a genuine forwarded identity: doing so would have to + // go through the interceptor, which replaces the whole set. + t.Run("a group appended to a stamped identity does not survive the interceptor", func(t *testing.T) { + asDaemon(t, root) + injected := metadata.MD{} + injected.Append(mdFwdGroup, sidAdministrators) + + ctx := WithForwardedIdentity(metadata.NewOutgoingContext(context.Background(), injected), + Identity{SID: "S-1-5-21-1-2-3-1001"}, true) + out, ok := metadata.FromOutgoingContext(ctx) + if !ok { + t.Fatal("no outgoing metadata") + } + if groups := out.Get(mdFwdGroup); len(groups) != 0 { + t.Fatalf("injected group survived: %v", groups) + } + }) +} + +func TestIsReservedForwardKey(t *testing.T) { + for _, key := range forwardKeys { + if !IsReservedForwardKey(key) { + t.Errorf("%q must be reserved", key) + } + } + + // grpc-gateway canonicalises header names, so the check has to be + // case-insensitive. + if !IsReservedForwardKey("X-Netbird-Fwd-Uid") { + t.Error("the check must be case-insensitive") + } + + for _, key := range []string{"authorization", "x-netbird", "x-netbird-fwd-uid-extra", ""} { + if IsReservedForwardKey(key) { + t.Errorf("%q must not be reserved", key) + } + } +} + +func TestForwardIdentityMetadata_AlwaysMarksForwarded(t *testing.T) { + for _, tc := range []struct { + name string + id Identity + known bool + }{ + {"known unix identity", unprivUser, true}, + {"unknown identity", Identity{}, false}, + {"windows identity", Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, true}, + } { + t.Run(tc.name, func(t *testing.T) { + md := ForwardIdentityMetadata(tc.id, tc.known) + if got := md.Get(mdFwd); len(got) != 1 || got[0] != "1" { + t.Fatalf("marker = %v, want exactly one \"1\"", got) + } + }) + } +} diff --git a/client/internal/ipcauth/identity.go b/client/internal/ipcauth/identity.go new file mode 100644 index 000000000..ff70c209a --- /dev/null +++ b/client/internal/ipcauth/identity.go @@ -0,0 +1,127 @@ +// Package ipcauth provides the kernel-authenticated identity of a local IPC +// (gRPC) caller and the transport credentials that surface it into the gRPC +// context, so the daemon can authorize individual RPCs by caller identity. +// +// On Unix the identity is read from the kernel via SO_PEERCRED (Linux) or +// LOCAL_PEERCRED (Darwin/FreeBSD). On Windows it is derived from the +// named-pipe client token. Platforms without a peer-identity primitive get no +// credentials, and every consumer must fail closed when no identity is +// available. +package ipcauth + +import ( + "context" + "fmt" + "slices" + + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" +) + +// Well-known Windows SIDs that identify a fully privileged principal. +const ( + sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM + sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE + sidNetworkService = "S-1-5-20" // NT AUTHORITY\NETWORK SERVICE + sidAdministrators = "S-1-5-32-544" // BUILTIN\Administrators +) + +// Identity is the kernel-authenticated identity of a local IPC caller. The +// zero value is not a valid identity: consumers must only use one obtained +// with a true ok/nil error return. +type Identity struct { + // UID and GID are the caller's Unix user ID and primary group ID. Both are + // zero on Windows, where SID is authoritative instead. + UID uint32 + GID uint32 + + // SID is the caller's Windows security identifier, empty on Unix. + SID string + + // Groups holds the caller's Windows group SIDs, captured from the client + // token at handshake time. Only groups that are enabled and not + // deny-only are captured, so a group listed here is one the caller can + // actually exercise. Empty on Unix. + Groups []string + + // Elevated reports whether the Windows client token is elevated (running + // as administrator, or an administrator with UAC turned off). Always false + // on Unix, where privilege is uid 0. + Elevated bool + + // PID is the caller's process ID where the platform reports it (Linux's + // SO_PEERCRED), and 0 where it does not. It identifies the daemon's own + // process dialling itself, which is what the JSON gateway does, and is never + // used to grant anything. + PID int32 +} + +// IsWindows reports whether this identity is a Windows principal (SID-based) +// rather than a Unix uid/gid principal. +func (i Identity) IsWindows() bool { + return i.SID != "" +} + +// IsPrivileged reports whether the caller is the platform's administrative +// principal, which is what the daemon requires for changes that cross the +// 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 +// 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 +// nothing about what this token may do on this machine. +func (i Identity) IsPrivileged() bool { + if !i.IsWindows() { + return i.UID == 0 + } + + if i.Elevated { + return true + } + + switch i.SID { + case sidLocalSystem, sidLocalService, sidNetworkService: + return true + } + + return slices.Contains(i.Groups, sidAdministrators) +} + +// String renders the identity for audit logs and denial messages. +func (i Identity) String() string { + if i.IsWindows() { + return fmt.Sprintf("sid=%s elevated=%t", i.SID, i.Elevated) + } + return fmt.Sprintf("uid=%d gid=%d", i.UID, i.GID) +} + +// AuthInfo carries the peer Identity as a gRPC credentials.AuthInfo so +// handlers can retrieve it from the request context via IdentityFromContext. +type AuthInfo struct { + credentials.CommonAuthInfo + Identity Identity +} + +// AuthType identifies the authentication scheme. +func (AuthInfo) AuthType() string { return "netbird-ipc-peercred" } + +// IdentityFromContext extracts the caller's kernel-authenticated identity from +// the gRPC peer context. The second return value is false when no IPC +// transport credentials were negotiated, which happens on a TCP daemon socket +// and on platforms without a peer-identity primitive. Callers MUST fail closed +// in that case. +func IdentityFromContext(ctx context.Context) (Identity, bool) { + p, ok := peer.FromContext(ctx) + if !ok { + return Identity{}, false + } + info, ok := p.AuthInfo.(AuthInfo) + if !ok { + return Identity{}, false + } + return info.Identity, true +} diff --git a/client/internal/ipcauth/ownedfile.go b/client/internal/ipcauth/ownedfile.go new file mode 100644 index 000000000..be7bf4864 --- /dev/null +++ b/client/internal/ipcauth/ownedfile.go @@ -0,0 +1,63 @@ +package ipcauth + +import ( + "fmt" + "os" +) + +// OpenOwnedFile opens path for reading on behalf of the IPC caller identified by +// id, and fails unless the opened file is a regular file that id owns. +// +// It exists for the paths a local caller hands to the daemon over the IPC. The +// daemon runs as root, so opening such a path unchecked lets any local user read +// any file through it. Ownership is the invariant that keeps the daemon from +// reading, with its own privileges, a file the caller could not read itself: a +// symlink or hard link planted at the path resolves to a file someone else owns +// and is refused. +// +// The check is made against the open descriptor rather than the path, so +// swapping the path between the check and the read cannot change the answer. +// +// A privileged caller is exempt: it can read the file directly, so refusing it +// here would protect nothing. The regular-file requirement still applies to +// everyone, since a fifo or device planted at the path is never a log file. +func OpenOwnedFile(id Identity, path string) (*os.File, error) { + f, err := openForRead(path) + if err != nil { + return nil, err + } + + if err := checkOwnership(id, f); err != nil { + if cerr := f.Close(); cerr != nil { + return nil, fmt.Errorf("%w (close: %v)", err, cerr) + } + return nil, err + } + + return f, nil +} + +func checkOwnership(id Identity, f *os.File) error { + info, err := f.Stat() + if err != nil { + return fmt.Errorf("stat %s: %w", f.Name(), err) + } + + if !info.Mode().IsRegular() { + return fmt.Errorf("%s is not a regular file", f.Name()) + } + + if IsPrivilegedCaller(id) { + return nil + } + + owned, err := fileOwnedBy(id, f) + if err != nil { + return fmt.Errorf("read owner of %s: %w", f.Name(), err) + } + if !owned { + return fmt.Errorf("%s is not owned by the caller (%s)", f.Name(), id) + } + + return nil +} diff --git a/client/internal/ipcauth/ownedfile_test.go b/client/internal/ipcauth/ownedfile_test.go new file mode 100644 index 000000000..b8178ad45 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_test.go @@ -0,0 +1,64 @@ +package ipcauth + +import ( + "io" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/require" +) + +// otherIdentity is an unprivileged caller that owns nothing the test creates. +func otherIdentity(t *testing.T) Identity { + t.Helper() + if runtime.GOOS == "windows" { + return Identity{SID: "S-1-5-21-1-2-3-1001"} + } + return Identity{UID: uint32(os.Geteuid() + 1), GID: uint32(os.Getegid() + 1)} +} + +func TestOpenOwnedFileReadsFileOwnedByCaller(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("hello"), 0600)) + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + f, err := OpenOwnedFile(id, path) + require.NoError(t, err) + t.Cleanup(func() { _ = f.Close() }) + + content, err := io.ReadAll(f) + require.NoError(t, err) + require.Equal(t, "hello", string(content)) +} + +func TestOpenOwnedFileRefusesFileOwnedByAnother(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("secret"), 0600)) + + _, err := OpenOwnedFile(otherIdentity(t), path) + require.ErrorContains(t, err, "not owned by the caller") +} + +func TestOpenOwnedFileRefusesNonRegularFile(t *testing.T) { + dir := t.TempDir() + + // The caller owns the directory, so this is the regular-file requirement + // talking, not the ownership check. + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, dir) + require.ErrorContains(t, err, "not a regular file") +} + +func TestOpenOwnedFileRefusesMissingFile(t *testing.T) { + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, filepath.Join(t.TempDir(), "absent.log")) + require.Error(t, err) +} diff --git a/client/internal/ipcauth/ownedfile_unix.go b/client/internal/ipcauth/ownedfile_unix.go new file mode 100644 index 000000000..4a6afcea7 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_unix.go @@ -0,0 +1,35 @@ +//go:build !windows + +package ipcauth + +import ( + "fmt" + "os" + "syscall" +) + +// openForRead opens a caller-supplied path without following a symlink at its +// final component and without blocking: a fifo planted at the path would +// otherwise stall the open until a writer appears, and the daemon holds a lock +// while it collects the file. +func openForRead(path string) (*os.File, error) { + f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_NONBLOCK, 0) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + return f, nil +} + +func fileOwnedBy(id Identity, f *os.File) (bool, error) { + info, err := f.Stat() + if err != nil { + return false, err + } + + stat, ok := info.Sys().(*syscall.Stat_t) + if !ok { + return false, fmt.Errorf("no owner information in %T", info.Sys()) + } + + return stat.Uid == id.UID, nil +} diff --git a/client/internal/ipcauth/ownedfile_unix_test.go b/client/internal/ipcauth/ownedfile_unix_test.go new file mode 100644 index 000000000..9e7831991 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_unix_test.go @@ -0,0 +1,57 @@ +//go:build !windows + +package ipcauth + +import ( + "errors" + "os" + "path/filepath" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// A symlink is the shape the arbitrary-read attempt takes: the caller owns the +// link, the file it points at belongs to someone else. +func TestOpenOwnedFileRefusesSymlink(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "target.log") + require.NoError(t, os.WriteFile(target, []byte("secret"), 0600)) + + link := filepath.Join(dir, "gui-client.log") + require.NoError(t, os.Symlink(target, link)) + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, link) + // O_NOFOLLOW on a symlink reports ELOOP on Linux/Darwin and EMLINK on FreeBSD. + if !errors.Is(err, syscall.ELOOP) && !errors.Is(err, syscall.EMLINK) { + t.Fatalf("symlink open: got %v, want ELOOP or EMLINK", err) + } +} + +// A fifo would block the open until a writer showed up, stalling the daemon +// while it holds its lock. +func TestOpenOwnedFileRefusesFifoWithoutBlocking(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, syscall.Mkfifo(path, 0600)) + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + done := make(chan error, 1) + go func() { + _, err := OpenOwnedFile(id, path) + done <- err + }() + + select { + case err := <-done: + require.ErrorContains(t, err, "not a regular file") + case <-time.After(5 * time.Second): + t.Fatal("opening a fifo blocked") + } +} diff --git a/client/internal/ipcauth/ownedfile_windows.go b/client/internal/ipcauth/ownedfile_windows.go new file mode 100644 index 000000000..19acaec4d --- /dev/null +++ b/client/internal/ipcauth/ownedfile_windows.go @@ -0,0 +1,59 @@ +//go:build windows + +package ipcauth + +import ( + "fmt" + "os" + + "golang.org/x/sys/windows" +) + +// openForRead opens a caller-supplied path without following a reparse point at +// it. FILE_FLAG_OPEN_REPARSE_POINT is the Windows analogue of O_NOFOLLOW: it +// opens a symlink/junction itself rather than its target, so the regular-file +// check in checkOwnership refuses a link the caller planted to redirect the +// read. FILE_FLAG_BACKUP_SEMANTICS lets a directory open too (as os.Open does), +// so a directory planted at the path is refused as non-regular rather than +// erroring here. The share mode matches os.Open so a log being written stays +// openable. +func openForRead(path string) (*os.File, error) { + p, err := windows.UTF16PtrFromString(path) + if err != nil { + return nil, fmt.Errorf("convert path %s: %w", path, err) + } + + handle, err := windows.CreateFile( + p, + windows.GENERIC_READ, + windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE, + nil, + windows.OPEN_EXISTING, + windows.FILE_FLAG_OPEN_REPARSE_POINT|windows.FILE_FLAG_BACKUP_SEMANTICS, + 0, + ) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + + return os.NewFile(uintptr(handle), path), nil +} + +// fileOwnedBy compares the file's owner SID with the caller's. Files an elevated +// process creates are owned by BUILTIN\Administrators rather than by the user, +// but such a caller is privileged and never reaches this check. +func fileOwnedBy(id Identity, f *os.File) (bool, error) { + // x/sys/windows GetSecurityInfo frees the OS buffer itself and returns a + // Go-heap copy, so there is nothing to LocalFree here. + sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION) + if err != nil { + return false, fmt.Errorf("read security info: %w", err) + } + + owner, _, err := sd.Owner() + if err != nil { + return false, fmt.Errorf("read owner: %w", err) + } + + return id.SID != "" && owner.String() == id.SID, nil +} diff --git a/client/internal/ipcauth/ownedfile_windows_test.go b/client/internal/ipcauth/ownedfile_windows_test.go new file mode 100644 index 000000000..ab68fdf39 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_windows_test.go @@ -0,0 +1,78 @@ +//go:build windows + +package ipcauth + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + "golang.org/x/sys/windows" +) + +// fileOwnerSID reads the owner SID of path the same way OpenOwnedFile does, so +// the test can construct an Identity that matches (or deliberately does not). +func fileOwnerSID(t *testing.T, path string) string { + t.Helper() + f, err := os.Open(path) + require.NoError(t, err) + t.Cleanup(func() { _ = f.Close() }) + + sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION) + require.NoError(t, err) + owner, _, err := sd.Owner() + require.NoError(t, err) + return owner.String() +} + +// The allow branch of fileOwnedBy is the SID-equality path the legitimate GUI +// flow depends on. Running elevated, a created file is owned by +// BUILTIN\Administrators; an Identity carrying that SID with Elevated=false and +// no groups is unprivileged by IsPrivileged (which reads the token, not the +// SID's RID), so this exercises the real GetSecurityInfo equality rather than +// the privileged-caller shortcut. +func TestOpenOwnedFileWindowsOwnerMatchAllows(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("hello"), 0600)) + + ownerSID := fileOwnerSID(t, path) + id := Identity{SID: ownerSID} + require.False(t, id.IsPrivileged(), "identity built from the owner SID must be unprivileged for this to test the match path") + + f, err := OpenOwnedFile(id, path) + require.NoError(t, err) + _ = f.Close() +} + +func TestOpenOwnedFileWindowsOwnerMismatchRefuses(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("secret"), 0600)) + + other := Identity{SID: "S-1-5-21-9-9-9-9999"} + require.False(t, other.IsPrivileged()) + + _, err := OpenOwnedFile(other, path) + require.ErrorContains(t, err, "not owned by the caller") +} + +// FILE_FLAG_OPEN_REPARSE_POINT must make OpenOwnedFile refuse a symlink the same +// way O_NOFOLLOW does on Unix, so a planted link can't redirect the read to +// another file. Creating a symlink needs a privilege the runner may lack, so the +// test skips rather than fails when it can't. +func TestOpenOwnedFileWindowsRefusesSymlink(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "target.log") + require.NoError(t, os.WriteFile(target, []byte("secret"), 0600)) + + link := filepath.Join(dir, "gui-client.log") + if err := os.Symlink(target, link); err != nil { + t.Skipf("cannot create symlink (privilege not held?): %v", err) + } + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, link) + require.Error(t, err, "a symlink must be refused") +} diff --git a/client/internal/ipcauth/peercred_bsd.go b/client/internal/ipcauth/peercred_bsd.go new file mode 100644 index 000000000..6d9c5247f --- /dev/null +++ b/client/internal/ipcauth/peercred_bsd.go @@ -0,0 +1,43 @@ +//go:build darwin || freebsd + +package ipcauth + +import ( + "fmt" + "net" + + "golang.org/x/sys/unix" +) + +// PeerIdentity reads the kernel-authenticated identity of the process on the +// other end of a Unix socket via LOCAL_PEERCRED. The xucred is recorded by the +// kernel at connect() time and carries the peer's uid and its group list, of +// which the first entry is the primary group. +func PeerIdentity(conn net.Conn) (Identity, error) { + uc, ok := conn.(*net.UnixConn) + if !ok { + return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn) + } + + raw, err := uc.SyscallConn() + if err != nil { + return Identity{}, fmt.Errorf("raw conn: %w", err) + } + + var cred *unix.Xucred + var credErr error + if err := raw.Control(func(fd uintptr) { + cred, credErr = unix.GetsockoptXucred(int(fd), unix.SOL_LOCAL, unix.LOCAL_PEERCRED) + }); err != nil { + return Identity{}, fmt.Errorf("control raw conn: %w", err) + } + if credErr != nil { + return Identity{}, fmt.Errorf("read LOCAL_PEERCRED: %w", credErr) + } + + id := Identity{UID: cred.Uid} + if cred.Ngroups > 0 { + id.GID = cred.Groups[0] + } + return id, nil +} diff --git a/client/internal/ipcauth/peercred_linux.go b/client/internal/ipcauth/peercred_linux.go new file mode 100644 index 000000000..417cc1e00 --- /dev/null +++ b/client/internal/ipcauth/peercred_linux.go @@ -0,0 +1,39 @@ +//go:build linux + +package ipcauth + +import ( + "fmt" + "net" + + "golang.org/x/sys/unix" +) + +// PeerIdentity reads the kernel-authenticated identity of the process on the +// other end of a Unix socket via SO_PEERCRED. The credentials are recorded by +// the kernel at connect() time and cannot be changed for the life of the +// connection, so they are not spoofable by the caller. +func PeerIdentity(conn net.Conn) (Identity, error) { + uc, ok := conn.(*net.UnixConn) + if !ok { + return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn) + } + + raw, err := uc.SyscallConn() + if err != nil { + return Identity{}, fmt.Errorf("raw conn: %w", err) + } + + var cred *unix.Ucred + var credErr error + if err := raw.Control(func(fd uintptr) { + cred, credErr = unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED) + }); err != nil { + return Identity{}, fmt.Errorf("control raw conn: %w", err) + } + if credErr != nil { + return Identity{}, fmt.Errorf("read SO_PEERCRED: %w", credErr) + } + + return Identity{UID: cred.Uid, GID: cred.Gid, PID: cred.Pid}, nil +} diff --git a/client/internal/ipcauth/pipeserver_windows.go b/client/internal/ipcauth/pipeserver_windows.go new file mode 100644 index 000000000..7ba59d574 --- /dev/null +++ b/client/internal/ipcauth/pipeserver_windows.go @@ -0,0 +1,87 @@ +//go:build windows + +package ipcauth + +import ( + "fmt" + "net" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" +) + +// PipeServerTrusted reports an error unless the pipe behind conn was created by a +// principal this client may hand secrets to. Clients call it for a pipe whose name +// carries no guarantee of its own, which is any name outside the +// ProtectedPrefix\Administrators namespace: that namespace already restricts +// creation to administrators and LocalSystem, while a plain name can be created by +// any local user before the daemon gets there. +// +// The decision is made from the pipe object's owner, not from the serving process, +// because a client cannot open a process running as another user at all, and the +// legitimate case is precisely an unprivileged client talking to a privileged +// daemon. Trusted owners are the service accounts, BUILTIN\Administrators, and +// this client's own user, the last of which is the daemon a user runs themselves +// as in netstack mode. A pipe owned by anyone else gets no setup key, pre-shared +// key or SSO prompt out of this client. +func PipeServerTrusted(conn net.Conn) error { + // go-winio's pipe connection embeds *win32File, which exposes Fd(). + fdConn, ok := conn.(interface{ Fd() uintptr }) + if !ok { + return fmt.Errorf("connection %T does not expose a pipe handle", conn) + } + + owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd())) + if err != nil { + return err + } + + if !trustedPipeOwner(owner) { + return fmt.Errorf("pipe owned by %s, which is neither an administrator nor this user", owner) + } + return nil +} + +// PipeOwnedBySelf reports whether the pipe behind conn was created by this very +// user, which is how a client recognises a daemon running as itself. Ownership it +// cannot read is reported as false. +func PipeOwnedBySelf(conn net.Conn) bool { + fdConn, ok := conn.(interface{ Fd() uintptr }) + if !ok { + return false + } + + owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd())) + if err != nil { + log.Debugf("read daemon pipe owner: %v", err) + return false + } + return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID +} + +// pipeOwnerSID reads the owner of the pipe object a client is connected to. The +// handle was opened with GENERIC_READ, which includes READ_CONTROL, so no extra +// access is needed. +func pipeOwnerSID(handle windows.Handle) (string, error) { + sd, err := windows.GetSecurityInfo(handle, windows.SE_KERNEL_OBJECT, windows.OWNER_SECURITY_INFORMATION) + if err != nil { + return "", fmt.Errorf("read pipe security info: %w", err) + } + + owner, _, err := sd.Owner() + if err != nil { + return "", fmt.Errorf("read pipe owner: %w", err) + } + return owner.String(), nil +} + +// trustedPipeOwner reports whether a pipe's owner is a principal a client may +// speak to. An elevated process's objects are owned by BUILTIN\Administrators by +// default, an unelevated one's by the user, which is why both forms appear here. +func trustedPipeOwner(owner string) bool { + switch owner { + case sidLocalSystem, sidLocalService, sidNetworkService, sidAdministrators: + return true + } + return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID +} diff --git a/client/internal/ipcauth/privileged.go b/client/internal/ipcauth/privileged.go new file mode 100644 index 000000000..95f2a50e9 --- /dev/null +++ b/client/internal/ipcauth/privileged.go @@ -0,0 +1,125 @@ +package ipcauth + +import ( + "os" + "runtime" +) + +// Fields of the ErrorInfo detail the daemon attaches to a PermissionDenied it +// raises for an operation that requires root/administrator. Clients match on +// Reason and Domain rather than on the message text, and render the summary and +// command themselves so the user gets guidance instead of a gRPC error dump. +const ( + // ErrorReasonPrivilegeRequired identifies the detail. + ErrorReasonPrivilegeRequired = "PRIVILEGE_REQUIRED" + // ErrorDomain scopes the reason to the NetBird daemon. + ErrorDomain = "daemon.netbird.io" + // ErrorMetaSummary is the one-sentence explanation of what was refused. + ErrorMetaSummary = "summary" + // ErrorMetaCommand is the command that performs the same operation with the + // privileges it needs, ready to copy and run. + ErrorMetaCommand = "command" +) + +// The identity of the process evaluating callers, captured once because it cannot +// change. selfKnown is false when it could not be read, in which case nothing is +// ever treated as this process. selfMayDelegate additionally requires this +// process to be unprivileged: see IsPrivilegedCaller. +var ( + selfIdentity Identity + selfKnown bool + selfMayDelegate bool + // selfPID is this process's PID, used to recognise the daemon dialling itself. + selfPID = os.Getpid() +) + +func init() { + id, err := CurrentProcessIdentity() + if err != nil { + return + } + selfIdentity, selfKnown = id, true + // Only an unprivileged daemon delegates its authority to its own identity. + // When it is root or LocalSystem, sharing its identity does not mean sharing + // its power: on Windows a filtered and a full token carry the same SID, so + // 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() +} + +// IsDaemonSelf reports whether an identity is this very process. The JSON gateway +// runs inside the daemon and re-dials it locally, so this is what distinguishes +// the gateway from any other caller, whatever user the daemon runs as. +func IsDaemonSelf(id Identity) bool { + if !selfKnown || id.IsWindows() != selfIdentity.IsWindows() { + return false + } + if id.IsWindows() { + return id.SID != "" && id.SID == selfIdentity.SID + } + return id.UID == selfIdentity.UID +} + +// IsPrivilegedCaller reports whether an identity may make the changes the daemon +// restricts to the platform administrator. This is the daemon's own rule and +// cannot be evaluated by a client, which does not know what the daemon runs as. +// +// Beyond root/administrator it accepts a caller running as the daemon's own +// identity when the daemon is itself unprivileged. That keeps a rootless container +// working, where there is no uid 0 at all, and a Windows daemon in netstack mode, +// which needs no administrator rights. In those setups a caller sharing the +// daemon's identity can already rewrite the config files it reads and replace the +// binary it runs, so refusing it a config change would protect nothing; and an +// unprivileged daemon cannot hand out a root shell in the first place. +func IsPrivilegedCaller(id Identity) bool { + if id.IsPrivileged() { + return true + } + return selfMayDelegate && IsDaemonSelf(id) +} + +// SelfDelegatesTo returns the identity this process delegates its authority to, +// and whether it delegates at all. Only an unprivileged daemon does: see +// IsPrivilegedCaller. It exists so a refusal can name who may actually perform the +// operation, because on such a host root is neither required nor necessarily +// available. +func SelfDelegatesTo() (Identity, bool) { + if !selfKnown || !selfMayDelegate { + return Identity{}, false + } + return selfIdentity, true +} + +// PrivilegedActor names the principal a privileged operation requires, for use +// in messages shown to the user. +func PrivilegedActor() string { + if runtime.GOOS == "windows" { + return "administrator privileges" + } + return "root" +} + +// ElevatedCommand renders a command so that running it grants the privileges the +// operation needs. Windows has no in-line equivalent of sudo, so the command is +// returned unchanged and the user is expected to run it from an elevated +// terminal. +func ElevatedCommand(command string) string { + if runtime.GOOS == "windows" { + return command + } + return "sudo " + command +} + +// UpCommand renders an elevated `netbird up` with the given flags, preceded by a +// `down`. The down is what makes the command work on a connected client: `netbird +// up` prints "Already connected" and returns without applying any config flag, so +// on its own the command would appear to do nothing. It is a no-op, exit 0, when +// the client is not connected. +// +// ";" rather than "&&" so the line can be pasted into any of the shells a user +// might have: PowerShell 5.1, still the default on Windows Server, rejects "&&" +// as a syntax error. +func UpCommand(flags string) string { + return ElevatedCommand("netbird down") + "; " + ElevatedCommand("netbird up "+flags) +} diff --git a/client/internal/ipcauth/privileged_test.go b/client/internal/ipcauth/privileged_test.go new file mode 100644 index 000000000..c1c7c1543 --- /dev/null +++ b/client/internal/ipcauth/privileged_test.go @@ -0,0 +1,134 @@ +package ipcauth + +import "testing" + +// The self rule is the one place privilege is granted to something other than the +// platform administrator, so its two guards matter: it must apply only when the +// daemon is itself unprivileged, and only to a caller with the daemon's identity. +func TestIsPrivilegedCaller_SelfRule(t *testing.T) { + tests := []struct { + name string + // self stands in for the process the daemon runs as. + self Identity + selfKnown bool + caller Identity + want bool + }{ + { + name: "root is privileged whatever the daemon runs as", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{UID: 0}, + want: true, + }, + { + name: "an unprivileged daemon delegates to its own user (rootless container)", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{UID: 1000}, + want: true, + }, + { + name: "an unprivileged daemon delegates to nobody else", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{UID: 1001}, + want: false, + }, + { + // The daemon is root on a normal install, so sharing its identity is + // already covered by being root; nothing else may match. + name: "a root daemon delegates to nobody", + self: Identity{UID: 0}, + selfKnown: true, + caller: Identity{UID: 1000}, + want: false, + }, + { + // Windows netstack mode: the daemon needs no administrator rights. + name: "an unprivileged windows daemon delegates to its own SID", + self: Identity{SID: "S-1-5-21-1-2-3-1001"}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-1001"}, + want: true, + }, + { + name: "an unprivileged windows daemon delegates to no other SID", + self: Identity{SID: "S-1-5-21-1-2-3-1001"}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-1002"}, + want: false, + }, + { + // The UAC boundary: a filtered and a full token of the same account + // carry the same SID but not the same power, so an elevated daemon must + // never delegate to its own SID. + name: "an elevated windows daemon does not delegate to its own SID", + self: Identity{SID: "S-1-5-21-1-2-3-500", Elevated: true}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-500"}, + want: false, + }, + { + name: "LocalSystem is privileged on its own merits, not by delegation", + self: Identity{SID: sidLocalSystem}, + selfKnown: true, + caller: Identity{SID: sidLocalSystem}, + want: true, // LocalSystem is privileged on its own merits + }, + { + name: "identities of different kinds never match", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-1001"}, + want: false, + }, + { + name: "an unknown self identity delegates to nobody", + self: Identity{}, + selfKnown: false, + caller: Identity{UID: 1000}, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate + t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) + + selfIdentity, selfKnown = tt.self, tt.selfKnown + selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged() + + if got := IsPrivilegedCaller(tt.caller); got != tt.want { + t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t", + tt.caller, tt.self, got, tt.want) + } + }) + } +} + +// The real process must never accidentally delegate: a test binary running as a +// normal user is unprivileged, so it may match itself, but nothing else. +func TestIsPrivilegedCaller_ThisProcess(t *testing.T) { + id, err := CurrentProcessIdentity() + if err != nil { + t.Skipf("cannot read this process's identity: %v", err) + } + + // This process is always allowed to act as itself: either it is privileged, or + // it is unprivileged and therefore delegates to its own identity. + if !IsPrivilegedCaller(id) { + t.Errorf("this process %v was refused its own identity", id) + } + + // A caller that is neither root nor this process must be refused, whatever + // this process happens to be. + other := Identity{UID: id.UID + 1} + if id.IsWindows() { + other = Identity{SID: id.SID + "9"} + } + if IsPrivilegedCaller(other) { + t.Errorf("an unrelated identity %v was treated as privileged", other) + } +} diff --git a/client/internal/ipcauth/self_unix.go b/client/internal/ipcauth/self_unix.go new file mode 100644 index 000000000..1b86c4fc0 --- /dev/null +++ b/client/internal/ipcauth/self_unix.go @@ -0,0 +1,17 @@ +//go:build !windows + +package ipcauth + +import "os" + +// CurrentProcessIdentity returns this process's identity as the daemon would +// see it if this process connected to the local IPC. It lets a client (the UI) +// decide up front whether a privileged operation can succeed, without a +// round-trip and without duplicating the rules: the answer comes from the same +// Identity.IsPrivileged the daemon applies. +func CurrentProcessIdentity() (Identity, error) { + return Identity{ + UID: uint32(os.Geteuid()), + GID: uint32(os.Getegid()), + }, nil +} diff --git a/client/internal/ipcauth/self_windows.go b/client/internal/ipcauth/self_windows.go new file mode 100644 index 000000000..5474cc101 --- /dev/null +++ b/client/internal/ipcauth/self_windows.go @@ -0,0 +1,35 @@ +//go:build windows + +package ipcauth + +import ( + "fmt" + + "golang.org/x/sys/windows" +) + +// CurrentProcessIdentity returns this process's identity as the daemon would see +// it if this process connected to the local IPC. It lets a client (the UI) +// decide up front whether a privileged operation can succeed, without a +// round-trip and without duplicating the rules: the answer comes from the same +// Identity.IsPrivileged the daemon applies to the token it reads off the pipe. +func CurrentProcessIdentity() (Identity, error) { + // A pseudo-token, so it must not be closed. + token := windows.GetCurrentProcessToken() + + user, err := token.GetTokenUser() + if err != nil { + return Identity{}, fmt.Errorf("read token user: %w", err) + } + + groups, err := tokenGroupSIDs(token) + if err != nil { + return Identity{}, err + } + + return Identity{ + SID: user.User.Sid.String(), + Groups: groups, + Elevated: token.IsElevated(), + }, nil +} diff --git a/client/internal/lazyconn/manager/manager.go b/client/internal/lazyconn/manager/manager.go index 3868e37e8..b7424bb2f 100644 --- a/client/internal/lazyconn/manager/manager.go +++ b/client/internal/lazyconn/manager/manager.go @@ -29,6 +29,11 @@ type managedPeer struct { type Config struct { InactivityThreshold *time.Duration + // ReconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is + // armed. The activity listener creates the wake peer with the overlay /32 only; without the + // routed prefixes WireGuard would not steer subnet-bound traffic to the wake endpoint, so an + // idle routing peer could never be woken by that traffic. Optional; nil disables the reconcile. + ReconcileAllowedIPs func(peerKey string) error } // Manager manages lazy connections @@ -56,6 +61,9 @@ type Manager struct { peerToHAGroups map[string][]route.HAUniqueID // peer ID -> HA groups they belong to haGroupToPeers map[route.HAUniqueID][]string // HA group -> peer IDs in the group routesMu sync.RWMutex + + // reconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is armed. + reconcileAllowedIPs func(peerKey string) error } // NewManager creates a new lazy connection manager @@ -73,6 +81,7 @@ func NewManager(config Config, engineCtx context.Context, peerStore *peerstore.S activityManager: activity.NewManager(wgIface), peerToHAGroups: make(map[string][]route.HAUniqueID), haGroupToPeers: make(map[route.HAUniqueID][]string), + reconcileAllowedIPs: config.ReconcileAllowedIPs, } if wgIface.IsUserspaceBind() { @@ -201,7 +210,7 @@ func (m *Manager) AddPeer(peerCfg lazyconn.PeerConfig) (bool, error) { return false, nil } - if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil { + if err := m.armActivityListener(peerCfg); err != nil { return false, err } @@ -288,7 +297,7 @@ func (m *Manager) DeactivatePeer(peerID peerid.ConnID) { m.inactivityManager.RemovePeer(mp.peerCfg.PublicKey) - if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil { + if err := m.armActivityListener(*mp.peerCfg); err != nil { mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err) return } @@ -465,6 +474,31 @@ func (m *Manager) close() { } // shouldDeferIdleForHA checks if peer should stay connected due to HA group requirements +// armRoutedAllowedIPs re-applies the peer's routed allowed IPs onto its freshly armed wake +// endpoint. The activity listener creates the wake peer with the overlay /32 only, so without +// this the routed prefixes would be missing and traffic to a routed subnet could not wake the +// idle routing peer. It is a no-op when no reconciler is configured. +// armActivityListener (re)arms the peer's wake endpoint via the activity manager and then +// re-applies its routed allowed IPs, so traffic to a routed subnet can wake an idle routing +// peer. The routed prefixes must be re-applied after the wake endpoint exists because the +// listener creates it with the overlay /32 only. +func (m *Manager) armActivityListener(peerCfg lazyconn.PeerConfig) error { + if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil { + return err + } + m.armRoutedAllowedIPs(&peerCfg) + return nil +} + +func (m *Manager) armRoutedAllowedIPs(peerCfg *lazyconn.PeerConfig) { + if m.reconcileAllowedIPs == nil { + return + } + if err := m.reconcileAllowedIPs(peerCfg.PublicKey); err != nil { + peerCfg.Log.Errorf("failed to reconcile routed allowed IPs on wake endpoint: %v", err) + } +} + func (m *Manager) shouldDeferIdleForHA(inactivePeers map[string]struct{}, peerID string) bool { m.routesMu.RLock() defer m.routesMu.RUnlock() @@ -577,7 +611,7 @@ func (m *Manager) onPeerInactivityTimedOut(peerIDs map[string]struct{}) { mp.peerCfg.Log.Infof("start activity monitor") - if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil { + if err := m.armActivityListener(*mp.peerCfg); err != nil { mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err) continue } diff --git a/client/internal/metrics/connection_type.go b/client/internal/metrics/connection_type.go index a3406a6b8..d393e5112 100644 --- a/client/internal/metrics/connection_type.go +++ b/client/internal/metrics/connection_type.go @@ -4,11 +4,17 @@ package metrics type ConnectionType string const ( - // ConnectionTypeICE represents a direct peer-to-peer connection using ICE - ConnectionTypeICE ConnectionType = "ice" + // ConnectionTypeICEP2P represents a direct peer-to-peer connection using ICE + ConnectionTypeICEP2P ConnectionType = "ice_p2p" + + // ConnectionTypeICETurn represents an ICE connection through a TURN server + ConnectionTypeICETurn ConnectionType = "ice_turn" // ConnectionTypeRelay represents a relayed connection ConnectionTypeRelay ConnectionType = "relay" + + // ConnectionTypeUnknown represents a connection with no active transport. It is not pushed. + ConnectionTypeUnknown ConnectionType = "unknown" ) // String returns the string representation of the connection type diff --git a/client/internal/metrics/influxdb_test.go b/client/internal/metrics/influxdb_test.go index b964e31a3..6a226fe2f 100644 --- a/client/internal/metrics/influxdb_test.go +++ b/client/internal/metrics/influxdb_test.go @@ -28,7 +28,7 @@ func TestInfluxDBMetrics_RecordAndExport(t *testing.T) { WgHandshakeSuccess: time.Now().Add(-1 * time.Second), } - m.RecordConnectionStages(context.Background(), agentInfo, "pair123", ConnectionTypeICE, false, ts) + m.RecordConnectionStages(context.Background(), agentInfo, "pair123", ConnectionTypeICEP2P, false, ts) var buf bytes.Buffer err := m.Export(&buf) @@ -60,7 +60,7 @@ func TestInfluxDBMetrics_ExportDeterministicFieldOrder(t *testing.T) { // Record multiple times and verify consistent field order for i := 0; i < 10; i++ { - m.RecordConnectionStages(context.Background(), agentInfo, "pair123", ConnectionTypeICE, false, ts) + m.RecordConnectionStages(context.Background(), agentInfo, "pair123", ConnectionTypeICEP2P, false, ts) } var buf bytes.Buffer diff --git a/client/internal/metrics/infra/README.md b/client/internal/metrics/infra/README.md index 7941a30cf..0a69404df 100644 --- a/client/internal/metrics/infra/README.md +++ b/client/internal/metrics/infra/README.md @@ -56,14 +56,33 @@ Measurement: `netbird_peer_connection` Tags: - `deployment_type`: "cloud" | "selfhosted" | "unknown" -- `connection_type`: "ice" | "relay" +- `connection_type`: "ice_p2p" | "ice_turn" | "relay" (see below) - `attempt_type`: "initial" | "reconnection" - `version`: NetBird version string - `os`: Operating system (linux, darwin, windows, android, ios, etc.) - `arch`: CPU architecture (amd64, arm64, etc.) +- `peer_id`: anonymised peer identifier (truncated SHA-256 of the WireGuard public key) +- `connection_pair_id`: deterministic identifier for the peer pair, identical on both sides **Note:** `SignalingReceived` is set when the first offer or answer arrives from the remote peer (in both initial and reconnection paths). It excludes the potentially unbounded wait for the remote peer to come online. +#### `connection_type` values + +Derived from the connection priority (`conntype.ConnPriority`) by `metricsConnType` in `client/internal/peer/conn.go`: + +| Value | Priority | Traffic is | +|-------|----------|------------| +| `ice_p2p` | `ICEP2P` | direct peer-to-peer | +| `ice_turn` | `ICETurn` | relayed, through a TURN server | +| `relay` | `Relay` | relayed, through a NetBird relay | +| `unknown` | `None` or unrecognised | no active transport — **the sample is not pushed** | + +**Direct traffic is `ice_p2p` only.** `ice_turn` is relayed despite being negotiated by ICE, matching `Conn.isRelayed`. + +`None` means no transport is active: not established yet, or reset after a relay drop or a peer-state reset. Such a sample cannot be attributed to a transport, so `recordConnectionMetrics` drops it instead of pushing it — `unknown` therefore never appears in the bucket. Connection counts are counts of connections whose transport was known at sampling time. + +**Samples recorded before 0.77 used a single `ice` value** which covered `ICEP2P`, `ICETurn` *and* `None`, so historical `ice` samples overstate direct connections by an unknown amount and must not be compared with `ice_p2p`. + ### Sync Duration Measurement: `netbird_sync` diff --git a/client/internal/mobile_dependency.go b/client/internal/mobile_dependency.go index 310d61a25..0234432b1 100644 --- a/client/internal/mobile_dependency.go +++ b/client/internal/mobile_dependency.go @@ -11,12 +11,14 @@ import ( // MobileDependency collect all dependencies for mobile platform type MobileDependency struct { - // Android only - TunAdapter device.TunAdapter - IFaceDiscover stdnet.ExternalIFaceDiscover + // Android and iOS NetworkChangeListener listener.NetworkChangeListener - HostDNSAddresses []netip.AddrPort - DnsReadyListener dns.ReadyListener + + // Android only + TunAdapter device.TunAdapter + IFaceDiscover stdnet.ExternalIFaceDiscover + HostDNSAddresses []netip.AddrPort + DnsReadyListener dns.ReadyListener // iOS only DnsManager dns.IosDnsManager diff --git a/client/internal/peer/conn.go b/client/internal/peer/conn.go index 09a4e8b02..83089606f 100644 --- a/client/internal/peer/conn.go +++ b/client/internal/peer/conn.go @@ -26,6 +26,7 @@ import ( "github.com/netbirdio/netbird/client/internal/portforward" "github.com/netbirdio/netbird/client/internal/rosenpass" "github.com/netbirdio/netbird/client/internal/stdnet" + "github.com/netbirdio/netbird/client/netevents" "github.com/netbirdio/netbird/route" relayClient "github.com/netbirdio/netbird/shared/relay/client" ) @@ -93,6 +94,10 @@ type ConnConfig struct { // ICEConfig ICE protocol configuration ICEConfig icemaker.Config + + // NetMgr gates the reconnection guard on OS-reported network + // availability; nil disables gating. + NetMgr *netevents.Manager } type Conn struct { @@ -254,7 +259,7 @@ func (conn *Conn) open(engineCtx context.Context, firstPacket []byte) error { conn.handshaker.AddICEListener(conn.workerICE.OnNewOffer) } - conn.guard = guard.NewGuard(conn.Log, conn.isConnectedOnAllWay, conn.config.Timeout, conn.srWatcher) + conn.guard = guard.NewGuard(conn.Log, conn.isConnectedOnAllWay, conn.config.Timeout, conn.srWatcher, conn.config.NetMgr) conn.wg.Add(1) go func() { @@ -307,6 +312,8 @@ func (conn *Conn) Close(signalToRemote bool) { if conn.wgWatcherCancel != nil { conn.wgWatcherCancel() + conn.wgWatcher = nil + conn.wgWatcherCancel = nil } conn.workerRelay.CloseConn() if conn.workerICE != nil { @@ -438,7 +445,7 @@ func (conn *Conn) onICEConnectionIsReady(priority conntype.ConnPriority, iceConn conn.dumpState.NewLocalProxy() wgProxy, err = conn.newProxy(iceConnInfo.RemoteConn) if err != nil { - conn.Log.Errorf("failed to add turn net.Conn to local proxy: %v", err) + conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err) return } ep = wgProxy.EndpointAddr() @@ -876,9 +883,8 @@ func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) { } wgProxy := conn.config.WgConfig.WgInterface.GetProxy() - if err := wgProxy.AddTurnConn(conn.ctx, udpAddr, remoteConn); err != nil { - conn.Log.Errorf("failed to add turn net.Conn to local proxy: %v", err) - return nil, err + if err := wgProxy.AddRelayedConn(conn.ctx, udpAddr, remoteConn); err != nil { + return nil, fmt.Errorf("add relayed conn to proxy: %w", err) } return wgProxy, nil } @@ -959,12 +965,9 @@ func (conn *Conn) recordConnectionMetrics() { priority := conn.currentConnPriority conn.mu.Unlock() - var connType metrics.ConnectionType - switch priority { - case conntype.Relay: - connType = metrics.ConnectionTypeRelay - default: - connType = metrics.ConnectionTypeICE + connType := metricsConnType(priority) + if connType == metrics.ConnectionTypeUnknown { + return } // Record metrics with timestamps - duration calculation happens in metrics package @@ -1065,3 +1068,16 @@ func boolToConnStatus(connected bool) guard.ConnStatus { } return guard.ConnStatusDisconnected } + +func metricsConnType(priority conntype.ConnPriority) metrics.ConnectionType { + switch priority { + case conntype.Relay: + return metrics.ConnectionTypeRelay + case conntype.ICETurn: + return metrics.ConnectionTypeICETurn + case conntype.ICEP2P: + return metrics.ConnectionTypeICEP2P + default: + return metrics.ConnectionTypeUnknown + } +} diff --git a/client/internal/peer/conn_test.go b/client/internal/peer/conn_test.go index 49979ea83..b709d5e40 100644 --- a/client/internal/peer/conn_test.go +++ b/client/internal/peer/conn_test.go @@ -11,6 +11,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/netbirdio/netbird/client/iface" + "github.com/netbirdio/netbird/client/internal/metrics" + "github.com/netbirdio/netbird/client/internal/peer/conntype" "github.com/netbirdio/netbird/client/internal/peer/dispatcher" "github.com/netbirdio/netbird/client/internal/peer/guard" "github.com/netbirdio/netbird/client/internal/peer/ice" @@ -386,3 +388,33 @@ func TestConn_onWGDisconnected_NoEscalationWithoutRosenpass(t *testing.T) { } assert.Empty(t, disconnected, "escalation must be limited to rosenpass connections") } + +func TestMetricsConnType(t *testing.T) { + tests := []struct { + name string + priority conntype.ConnPriority + expected metrics.ConnectionType + }{ + {"relay", conntype.Relay, metrics.ConnectionTypeRelay}, + {"ice over turn is relayed, not p2p", conntype.ICETurn, metrics.ConnectionTypeICETurn}, + {"direct p2p", conntype.ICEP2P, metrics.ConnectionTypeICEP2P}, + {"unset priority is unknown, not p2p", conntype.None, metrics.ConnectionTypeUnknown}, + {"unrecognised priority is unknown", conntype.ConnPriority(99), metrics.ConnectionTypeUnknown}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + assert.Equal(t, tc.expected, metricsConnType(tc.priority)) + }) + } +} + +func TestMetricsConnType_RelayedMatchesIsRelayed(t *testing.T) { + for _, priority := range []conntype.ConnPriority{conntype.None, conntype.Relay, conntype.ICETurn, conntype.ICEP2P} { + conn := &Conn{currentConnPriority: priority} + tag := metricsConnType(priority) + relayedTag := tag == metrics.ConnectionTypeRelay || tag == metrics.ConnectionTypeICETurn + assert.Equal(t, conn.isRelayed(), relayedTag, + "priority %s: isRelayed and the %q metric tag must agree", priority, tag) + } +} diff --git a/client/internal/peer/guard/guard.go b/client/internal/peer/guard/guard.go index 6c2e846a9..73bab2a89 100644 --- a/client/internal/peer/guard/guard.go +++ b/client/internal/peer/guard/guard.go @@ -22,6 +22,12 @@ const ( type connStatusFunc func() ConnStatus +// NetworkWatcher is the availability view the guard gates reconnects on. +type NetworkWatcher interface { + IsOnline() bool + Changed() <-chan struct{} +} + // Guard is responsible for the reconnection logic. // It will trigger to send an offer to the peer then has connection issues. // Watch these events: @@ -31,20 +37,26 @@ type connStatusFunc func() ConnStatus // - Relayed connection disconnected // - ICE candidate changes type Guard struct { - log *log.Entry - isConnectedOnAllWay connStatusFunc - timeout time.Duration - srWatcher *SRWatcher + log *log.Entry + isConnectedOnAllWay connStatusFunc + timeout time.Duration + srWatcher *SRWatcher + // netWatcher gates reconnect attempts on OS-reported network availability; + // nil disables gating. + netWatcher NetworkWatcher relayedConnDisconnected chan struct{} iCEConnDisconnected chan struct{} } -func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher) *Guard { +// NewGuard creates a reconnection guard for a peer connection. A nil netWatcher +// disables network availability gating. +func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher, netWatcher NetworkWatcher) *Guard { return &Guard{ log: log, isConnectedOnAllWay: isConnectedFn, timeout: timeout, srWatcher: srWatcher, + netWatcher: netWatcher, relayedConnDisconnected: make(chan struct{}, 1), iCEConnDisconnected: make(chan struct{}, 1), } @@ -96,9 +108,19 @@ func (g *Guard) reconnectLoopWithRetry(ctx context.Context, callback func()) { iceState := &iceRetryState{log: g.log} defer iceState.reset() + var netChanged <-chan struct{} + if g.netWatcher != nil { + netChanged = g.netWatcher.Changed() + } + for { select { case <-tickerChannel: + // skip attempts while the OS reports no usable network; the + // netChanged case below resumes the loop once it returns + if g.netWatcher != nil && !g.netWatcher.IsOnline() { + continue + } switch g.isConnectedOnAllWay() { case ConnStatusConnected: // all good, nothing to do @@ -135,6 +157,23 @@ func (g *Guard) reconnectLoopWithRetry(ctx context.Context, callback func()) { tickerChannel = ticker.C iceState.reset() + case <-netChanged: + // Re-arm for the next transition before acting on this one. + netChanged = g.netWatcher.Changed() + if !g.netWatcher.IsOnline() { + continue + } + // Ticks skipped while offline drove the backoff towards its + // maximum without ever attempting, and left the ICE budget + // frozen — possibly in hourly mode. Recover on our own so the + // peer does not depend on a signal or relay event that never + // comes when both stayed up across the outage. + g.log.Debugf("network is back, reset reconnection ticker") + ticker.Stop() + ticker = g.newReconnectTicker(ctx) + tickerChannel = ticker.C + iceState.reset() + case <-ctx.Done(): g.log.Debugf("context is done, stop reconnect loop") return diff --git a/client/internal/peer/guard/guard_leak_test.go b/client/internal/peer/guard/guard_leak_test.go index ded3e4aea..3d82ec591 100644 --- a/client/internal/peer/guard/guard_leak_test.go +++ b/client/internal/peer/guard/guard_leak_test.go @@ -15,7 +15,7 @@ import ( func newTestGuard(status connStatusFunc) *Guard { srw := NewSRWatcher(nil, nil, nil, ice.Config{}) - return NewGuard(log.WithField("test", "guard"), status, 50*time.Millisecond, srw) + return NewGuard(log.WithField("test", "guard"), status, 50*time.Millisecond, srw, nil) } // countBackoffTickerGoroutines returns how many goroutines are currently sitting diff --git a/client/internal/peer/guard/guard_netstate_test.go b/client/internal/peer/guard/guard_netstate_test.go new file mode 100644 index 000000000..44999cae1 --- /dev/null +++ b/client/internal/peer/guard/guard_netstate_test.go @@ -0,0 +1,107 @@ +package guard + +import ( + "context" + "sync/atomic" + "testing" + "time" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/peer/ice" + "github.com/netbirdio/netbird/client/netevents/netstate" +) + +// newTestGuardWithNetState builds a guard with a realistic MaxInterval: the +// backoff must be able to grow well past the outage, as it does in production +// where the timeout is seconds to minutes. +func newTestGuardWithNetState(status connStatusFunc, netState *netstate.State) *Guard { + srw := NewSRWatcher(nil, nil, nil, ice.Config{}) + return NewGuard(log.WithField("test", "guard"), status, 30*time.Second, srw, netState) +} + +// TestGuard_RecoversAfterOfflineToOnline covers a peer that stays disconnected +// across a network outage while neither signal nor relay reports an event — +// both stayed up, as on a short airplane mode toggle over Wi-Fi. +// +// Every tick taken while offline is skipped, but it still advances the +// exponential backoff, so by the time the network returns the next tick can be +// tens of seconds away. Without an explicit reaction to the transition the +// peer waits out that interval for a recovery that could start immediately. +func TestGuard_RecoversAfterOfflineToOnline(t *testing.T) { + netState := netstate.New() + + var attempts atomic.Int32 + g := newTestGuardWithNetState(func() ConnStatus { return ConnStatusDisconnected }, netState) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + // Start from the reconnect ticker (800ms initial interval), the state a + // peer is in after it loses its connection. + go g.Start(ctx, func() { attempts.Add(1) }) + g.SetRelayedConnDisconnected() + + // Let the backoff climb: 0.8s, 1.6s, 3.2s, 6.4s ... every tick is skipped + // while offline, but each one doubles the wait for the next. + netState.Set(false) + time.Sleep(8 * time.Second) + + offlineAttempts := attempts.Load() + if offlineAttempts != 0 { + t.Fatalf("callback ran %d times while offline, want 0", offlineAttempts) + } + + netState.Set(true) + + // The next organic tick is now several seconds out, so anything within + // this window can only come from reacting to the transition itself. + pollCtx, stopPolling := context.WithTimeout(ctx, 2*time.Second) + defer stopPolling() + + select { + case <-pollCtx.Done(): + t.Fatal("peer was not retried within 2s of the network coming back, " + + "with neither a signal nor a relay event to fall back on") + case <-pollUntil(pollCtx, func() bool { return attempts.Load() > 0 }): + } +} + +// TestGuard_OfflineTransitionDoesNotRetry checks the other direction: going +// offline must not itself trigger an attempt. +func TestGuard_OfflineTransitionDoesNotRetry(t *testing.T) { + netState := netstate.New() + + var attempts atomic.Int32 + g := newTestGuardWithNetState(func() ConnStatus { return ConnStatusDisconnected }, netState) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go g.Start(ctx, func() { attempts.Add(1) }) + + netState.Set(false) + time.Sleep(5 * time.Second) + + if got := attempts.Load(); got != 0 { + t.Fatalf("callback ran %d times after going offline, want 0", got) + } +} + +// pollUntil closes the returned channel once cond holds. It gives up when ctx +// is done, so the polling goroutine never outlives the test that started it. +func pollUntil(ctx context.Context, cond func() bool) <-chan struct{} { + done := make(chan struct{}) + go func() { + for { + if cond() { + close(done) + return + } + select { + case <-ctx.Done(): + return + case <-time.After(10 * time.Millisecond): + } + } + }() + return done +} diff --git a/client/internal/peer/handshaker.go b/client/internal/peer/handshaker.go index 56e82e6e3..6ecb2a947 100644 --- a/client/internal/peer/handshaker.go +++ b/client/internal/peer/handshaker.go @@ -81,14 +81,19 @@ type Handshaker struct { func NewHandshaker(log *log.Entry, config ConnConfig, signaler *Signaler, ice *WorkerICE, relay *WorkerRelay, metricsStages *MetricsStages) *Handshaker { h := &Handshaker{ - log: log, - config: config, - signaler: signaler, - ice: ice, - relay: relay, - metricsStages: metricsStages, - remoteOffersCh: make(chan OfferAnswer), - remoteAnswerCh: make(chan OfferAnswer), + log: log, + config: config, + signaler: signaler, + ice: ice, + relay: relay, + metricsStages: metricsStages, + // Buffered by one so an offer or answer that arrives between Open launching + // the Listen goroutine and it reaching its receive is held rather than + // dropped. A peer activated by an incoming signal receives the remote's + // message in that window; an unbuffered channel skips it as "receiver not + // ready", and the connection cannot proceed until the remote re-sends. + remoteOffersCh: make(chan OfferAnswer, 1), + remoteAnswerCh: make(chan OfferAnswer, 1), } // assume remote supports ICE until we learn otherwise from received offers h.remoteICESupported.Store(ice != nil) @@ -162,29 +167,38 @@ func (h *Handshaker) SendOffer() error { return h.sendOffer() } -// OnRemoteOffer handles an offer from the remote peer and returns true if the message was accepted, false otherwise -// doesn't block, discards the message if connection wasn't ready +// OnRemoteOffer hands an offer to Listen without blocking, keeping only the most +// recent one if several arrive before Listen reads them. func (h *Handshaker) OnRemoteOffer(offer OfferAnswer) { - select { - case h.remoteOffersCh <- offer: - return - default: - h.log.Warnf("skipping remote offer message because receiver not ready") - // connection might not be ready yet to receive so we ignore the message - return - } + enqueueLatest(h.remoteOffersCh, offer) } -// OnRemoteAnswer handles an offer from the remote peer and returns true if the message was accepted, false otherwise -// doesn't block, discards the message if connection wasn't ready +// OnRemoteAnswer hands an answer to Listen without blocking, keeping only the most +// recent one if several arrive before Listen reads them. func (h *Handshaker) OnRemoteAnswer(answer OfferAnswer) { + enqueueLatest(h.remoteAnswerCh, answer) +} + +// enqueueLatest delivers msg on a one-slot channel without blocking. When the slot +// already holds an unread message the older one is discarded in favor of msg, so a +// message arriving before Listen starts reading is held rather than dropped, and +// the newest wins if several arrive first. Safe because there is a single producer +// (the engine loop): after draining the stale value the send always has room. +func enqueueLatest(ch chan OfferAnswer, msg OfferAnswer) { select { - case h.remoteAnswerCh <- answer: + case ch <- msg: return default: - // connection might not be ready yet to receive so we ignore the message - h.log.Warnf("skipping remote answer message because receiver not ready") - return + } + + select { + case <-ch: + default: + } + + select { + case ch <- msg: + default: } } diff --git a/client/internal/peer/handshaker_test.go b/client/internal/peer/handshaker_test.go new file mode 100644 index 000000000..5e203d78b --- /dev/null +++ b/client/internal/peer/handshaker_test.go @@ -0,0 +1,63 @@ +package peer + +import ( + "testing" + "time" + + log "github.com/sirupsen/logrus" + "github.com/stretchr/testify/assert" +) + +func newTestHandshaker(t *testing.T) *Handshaker { + t.Helper() + // The tests exercise the answer path, whose Listen branch dispatches to the + // relay listener without sending an answer, so no signaler/ICE/relay is needed. + return NewHandshaker(log.WithField("test", t.Name()), ConnConfig{}, nil, nil, nil, nil) +} + +// TestHandshakerHoldsSignalArrivingBeforeListen covers the case where a peer is +// activated by an incoming signal: the remote's offer/answer arrives in the same +// step that opens the connection, before the Listen loop starts reading. The +// message must be held rather than dropped, or the connection cannot proceed until +// the remote re-sends. This is the path taken when an eager peer connects to a +// lazily-managed one. +func TestHandshakerHoldsSignalArrivingBeforeListen(t *testing.T) { + h := newTestHandshaker(t) + + processed := make(chan *OfferAnswer, 4) + h.AddRelayListener(func(o *OfferAnswer) { processed <- o }) + + // Delivered before Listen is reading, as when the peer is woken by the remote's + // signal and the message is delivered right after Open. + h.OnRemoteAnswer(OfferAnswer{WgListenPort: 51820}) + + go h.Listen(t.Context()) + + select { + case <-processed: + case <-time.After(2 * time.Second): + assert.Fail(t, "remote-answer dispatch: signal delivered before Listen was ready was dropped") + } +} + +// TestHandshakerKeepsLatestSignalBeforeListen covers several signals arriving +// before Listen reads: the newest must win (matching the latest-offer contract), +// rather than the first being kept and later ones discarded. +func TestHandshakerKeepsLatestSignalBeforeListen(t *testing.T) { + h := newTestHandshaker(t) + + processed := make(chan *OfferAnswer, 4) + h.AddRelayListener(func(o *OfferAnswer) { processed <- o }) + + h.OnRemoteAnswer(OfferAnswer{WgListenPort: 1111}) + h.OnRemoteAnswer(OfferAnswer{WgListenPort: 2222}) + + go h.Listen(t.Context()) + + select { + case got := <-processed: + assert.Equal(t, 2222, got.WgListenPort, "remote-answer dispatch: the latest queued signal should be processed") + case <-time.After(2 * time.Second): + assert.Fail(t, "remote-answer dispatch: queued signal was dropped") + } +} diff --git a/client/internal/peer/listener.go b/client/internal/peer/listener.go index c601fe534..2bb7fcf32 100644 --- a/client/internal/peer/listener.go +++ b/client/internal/peer/listener.go @@ -1,11 +1,40 @@ package peer +// ClientState identifies the client connection state delivered via +// Listener.OnStateChanged. +type ClientState int + +// Client states. The numeric values cross the gomobile boundary (the mobile +// bindings re-export them as integer constants), so they are a wire format: +// append new states at the end, never reorder or insert. +const ( + ClientStateDisconnected ClientState = iota + ClientStateConnected + ClientStateConnecting + ClientStateDisconnecting + // ClientStateNoNetwork is an overlay state: it is never stored as the + // last notification, only derived from ClientStateConnecting while the + // OS reports no usable network (see notifier.effectiveState). + ClientStateNoNetwork +) + // Listener is a callback type about the NetBird network connection state type Listener interface { + // OnStateChanged reports every client state transition. New states are + // delivered only through this callback; the per-state callbacks below + // are kept for compatibility and will be removed once all consumers + // have migrated. + OnStateChanged(state ClientState) + + // Deprecated: consume OnStateChanged instead. OnConnected() + // Deprecated: consume OnStateChanged instead. OnDisconnected() + // Deprecated: consume OnStateChanged instead. OnConnecting() + // Deprecated: consume OnStateChanged instead. OnDisconnecting() + OnAddressChanged(string, string) OnPeersListChanged(int) } diff --git a/client/internal/peer/notifier.go b/client/internal/peer/notifier.go index 8d1954fe5..1ee1d32ea 100644 --- a/client/internal/peer/notifier.go +++ b/client/internal/peer/notifier.go @@ -4,31 +4,64 @@ import ( "sync" ) -const ( - stateDisconnected = iota - stateConnected - stateConnecting - stateDisconnecting -) - type notifier struct { + // publishLock orders state publication: it is held across computing the + // effective state and handing it to the listener, so a transition cannot + // overtake a newer one and leave the listener on a stale state. + publishLock sync.Mutex serverStateLock sync.Mutex listenersLock sync.Mutex listener Listener currentClientState bool - lastNotification int + lastNotification ClientState lastNumberOfPeers int lastFqdnAddress string lastIPAddress string + networkAvailable bool } func newNotifier() *notifier { - return ¬ifier{} + return ¬ifier{ + networkAvailable: true, + } +} + +// effectiveState maps the computed state to what listeners should see: +// while the OS reports no usable network, "Connecting" would be a lie — +// connection attempts are suspended — so it is reported as NoNetwork. +// Caller must hold serverStateLock. +func (n *notifier) effectiveState(state ClientState) ClientState { + if !n.networkAvailable && state == ClientStateConnecting { + return ClientStateNoNetwork + } + return state +} + +// setNetworkAvailable records the OS network availability and re-notifies +// the listener when the flag flips the effective state (Connecting <-> +// NoNetwork). +func (n *notifier) setNetworkAvailable(available bool) { + n.publishLock.Lock() + defer n.publishLock.Unlock() + + n.serverStateLock.Lock() + if n.networkAvailable == available { + n.serverStateLock.Unlock() + return + } + previous := n.effectiveState(n.lastNotification) + n.networkAvailable = available + current := n.effectiveState(n.lastNotification) + n.serverStateLock.Unlock() + + if previous != current { + n.notify(current) + } } func (n *notifier) setListener(listener Listener) { n.serverStateLock.Lock() - lastNotification := n.lastNotification + lastNotification := n.effectiveState(n.lastNotification) numOfPeers := n.lastNumberOfPeers fqdnAddress := n.lastFqdnAddress address := n.lastIPAddress @@ -52,6 +85,9 @@ func (n *notifier) removeListener() { } func (n *notifier) updateServerStates(mgmState bool, signalState bool) { + n.publishLock.Lock() + defer n.publishLock.Unlock() + n.serverStateLock.Lock() calculatedState := n.calculateState(mgmState, signalState) @@ -61,43 +97,54 @@ func (n *notifier) updateServerStates(mgmState bool, signalState bool) { } n.lastNotification = calculatedState + effective := n.effectiveState(calculatedState) n.serverStateLock.Unlock() - n.notify(calculatedState) + n.notify(effective) } func (n *notifier) clientStart() { + n.publishLock.Lock() + defer n.publishLock.Unlock() + n.serverStateLock.Lock() n.currentClientState = true - n.lastNotification = stateConnecting + n.lastNotification = ClientStateConnecting + effective := n.effectiveState(ClientStateConnecting) n.serverStateLock.Unlock() - n.notify(stateConnecting) + n.notify(effective) } func (n *notifier) clientStop() { + n.publishLock.Lock() + defer n.publishLock.Unlock() + n.serverStateLock.Lock() n.currentClientState = false - n.lastNotification = stateDisconnected + n.lastNotification = ClientStateDisconnected n.serverStateLock.Unlock() - n.notify(stateDisconnected) + n.notify(ClientStateDisconnected) } func (n *notifier) clientTearDown() { + n.publishLock.Lock() + defer n.publishLock.Unlock() + n.serverStateLock.Lock() n.currentClientState = false - n.lastNotification = stateDisconnecting + n.lastNotification = ClientStateDisconnecting n.serverStateLock.Unlock() - n.notify(stateDisconnecting) + n.notify(ClientStateDisconnecting) } -func (n *notifier) isServerStateChanged(newState int) bool { +func (n *notifier) isServerStateChanged(newState ClientState) bool { return n.lastNotification != newState } -func (n *notifier) notify(state int) { +func (n *notifier) notify(state ClientState) { n.listenersLock.Lock() listener := n.listener n.listenersLock.Unlock() @@ -109,20 +156,20 @@ func (n *notifier) notify(state int) { notifyListener(listener, state) } -func (n *notifier) calculateState(managementConn, signalConn bool) int { +func (n *notifier) calculateState(managementConn, signalConn bool) ClientState { if managementConn && signalConn { - return stateConnected + return ClientStateConnected } if !managementConn && !signalConn && !n.currentClientState { - return stateDisconnected + return ClientStateDisconnected } - if n.lastNotification == stateDisconnecting { - return stateDisconnecting + if n.lastNotification == ClientStateDisconnecting { + return ClientStateDisconnecting } - return stateConnecting + return ClientStateConnecting } func (n *notifier) peerListChanged(numOfPeers int) { @@ -159,15 +206,19 @@ func (n *notifier) localAddressChanged(fqdn, address string) { listener.OnAddressChanged(fqdn, address) } -func notifyListener(l Listener, state int) { +func notifyListener(l Listener, state ClientState) { + // legacy per-state callbacks; NoNetwork is delivered only via + // OnStateChanged below switch state { - case stateDisconnected: + case ClientStateDisconnected: l.OnDisconnected() - case stateConnected: + case ClientStateConnected: l.OnConnected() - case stateConnecting: + case ClientStateConnecting: l.OnConnecting() - case stateDisconnecting: + case ClientStateDisconnecting: l.OnDisconnecting() } + + l.OnStateChanged(state) } diff --git a/client/internal/peer/notifier_concurrent_test.go b/client/internal/peer/notifier_concurrent_test.go new file mode 100644 index 000000000..fcaaaad3b --- /dev/null +++ b/client/internal/peer/notifier_concurrent_test.go @@ -0,0 +1,108 @@ +package peer + +import ( + "sync" + "testing" + "time" +) + +type recordingListener struct { + mu sync.Mutex + states []ClientState + onState func(ClientState) +} + +func (l *recordingListener) OnStateChanged(state ClientState) { + l.mu.Lock() + l.states = append(l.states, state) + hook := l.onState + l.mu.Unlock() + + if hook != nil { + hook(state) + } +} + +func (l *recordingListener) last() (ClientState, bool) { + l.mu.Lock() + defer l.mu.Unlock() + if len(l.states) == 0 { + return 0, false + } + return l.states[len(l.states)-1], true +} + +func (l *recordingListener) snapshot() []ClientState { + l.mu.Lock() + defer l.mu.Unlock() + return append([]ClientState(nil), l.states...) +} + +func (l *recordingListener) OnConnected() {} +func (l *recordingListener) OnDisconnected() {} +func (l *recordingListener) OnConnecting() {} +func (l *recordingListener) OnDisconnecting() {} +func (l *recordingListener) OnAddressChanged(string, string) {} +func (l *recordingListener) OnPeersListChanged(int) {} + +// TestNotifier_ConcurrentAvailabilityFlipOrdersPublication holds the first +// transition inside the listener callback and flips availability again from +// another goroutine while it is parked. The second flip must not publish +// ahead of the one in flight, otherwise the listener ends up on a state the +// notifier already superseded. +func TestNotifier_ConcurrentAvailabilityFlipOrdersPublication(t *testing.T) { + n := newNotifier() + n.currentClientState = true + n.lastNotification = ClientStateConnecting + + entered := make(chan struct{}) + release := make(chan struct{}) + + l := &recordingListener{} + l.onState = func(state ClientState) { + if state != ClientStateNoNetwork { + return + } + l.mu.Lock() + l.onState = nil + l.mu.Unlock() + close(entered) + <-release + } + n.listener = l + + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + n.setNetworkAvailable(false) + }() + + <-entered + + flipped := make(chan struct{}) + go func() { + defer close(flipped) + n.setNetworkAvailable(true) + }() + + select { + case <-flipped: + t.Fatal("the online transition published while the offline one was " + + "still in flight; publication is not serialized") + case <-time.After(200 * time.Millisecond): + } + + close(release) + <-flipped + wg.Wait() + + got, ok := l.last() + if !ok { + t.Fatal("listener never observed a state") + } + if got != ClientStateConnecting { + t.Fatalf("listener holds %v after the network came back, want Connecting; sequence: %v", + got, l.snapshot()) + } +} diff --git a/client/internal/peer/notifier_test.go b/client/internal/peer/notifier_test.go index 0b7722b0c..a73016b05 100644 --- a/client/internal/peer/notifier_test.go +++ b/client/internal/peer/notifier_test.go @@ -6,29 +6,32 @@ import ( ) type mocListener struct { - lastState int + lastState ClientState wg sync.WaitGroup peersWg sync.WaitGroup peers int } func (l *mocListener) OnConnected() { - l.lastState = stateConnected + l.lastState = ClientStateConnected l.wg.Done() } func (l *mocListener) OnDisconnected() { - l.lastState = stateDisconnected + l.lastState = ClientStateDisconnected l.wg.Done() } func (l *mocListener) OnConnecting() { - l.lastState = stateConnecting + l.lastState = ClientStateConnecting l.wg.Done() } func (l *mocListener) OnDisconnecting() { - l.lastState = stateDisconnecting + l.lastState = ClientStateDisconnecting l.wg.Done() } +func (l *mocListener) OnStateChanged(state ClientState) { + +} func (l *mocListener) OnAddressChanged(host, addr string) { } @@ -57,15 +60,15 @@ func Test_notifier_serverState(t *testing.T) { type scenario struct { name string - expected int + expected ClientState mgmState bool signalState bool } scenarios := []scenario{ - {"connected", stateConnected, true, true}, - {"mgm down", stateConnecting, false, true}, - {"signal down", stateConnecting, true, false}, - {"disconnected", stateDisconnected, false, false}, + {"connected", ClientStateConnected, true, true}, + {"mgm down", ClientStateConnecting, false, true}, + {"signal down", ClientStateConnecting, true, false}, + {"disconnected", ClientStateDisconnected, false, false}, } for _, tt := range scenarios { @@ -85,7 +88,7 @@ func Test_notifier_SetListener(t *testing.T) { listener.setPeersWaiter() n := newNotifier() - n.lastNotification = stateConnecting + n.lastNotification = ClientStateConnecting n.setListener(listener) listener.wait() listener.waitPeers() @@ -99,7 +102,7 @@ func Test_notifier_RemoveListener(t *testing.T) { listener.setWaiter() listener.setPeersWaiter() n := newNotifier() - n.lastNotification = stateConnecting + n.lastNotification = ClientStateConnecting n.setListener(listener) // setListener replays cached state on a goroutine; wait for both the state // and peers callbacks to finish so we don't race on listener.peers. diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index a987482fe..24e3e7fac 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -813,19 +813,14 @@ func (d *Status) SetSessionExpiresAt(deadline time.Time) { } // GetSessionExpiresAt returns the most recently recorded SSO session deadline, -// or the zero value when no deadline is tracked. A deadline that has already -// slipped into the past reports as "none": once the session has expired it is -// no longer a meaningful countdown, and the sessionwatch.Watcher does not -// arm a timer at the deadline itself to clear it (only the two pre-expiry -// warnings). Without this guard the UI would keep painting a stale -// "expires in …" against a moment that has passed until the next login, -// extend, or teardown rewrote the value. +// or the zero value when no deadline is tracked. A deadline in the past is +// returned as-is: it means the session has expired, and consumers (tray row, +// CLI status) render it as "expired" rather than hiding it — masking it as +// "none" would blank the UI at the exact moment it should say the session +// ended. func (d *Status) GetSessionExpiresAt() time.Time { d.mux.Lock() defer d.mux.Unlock() - if !d.sessionExpiresAt.IsZero() && d.sessionExpiresAt.Before(time.Now()) { - return time.Time{} - } return d.sessionExpiresAt } @@ -1216,6 +1211,12 @@ func (d *Status) ClientTeardown() { d.notifyStateChange() } +// SetNetworkAvailable records the OS-reported network availability; while +// unavailable, listeners see NoNetwork instead of Connecting. +func (d *Status) SetNetworkAvailable(available bool) { + d.notifier.setNetworkAvailable(available) +} + // SetConnectionListener set a listener to the notifier func (d *Status) SetConnectionListener(listener Listener) { d.notifier.setListener(listener) diff --git a/client/internal/peer/worker_ice.go b/client/internal/peer/worker_ice.go index b1aa3e0f9..83cac13f5 100644 --- a/client/internal/peer/worker_ice.go +++ b/client/internal/peer/worker_ice.go @@ -255,8 +255,8 @@ func (w *WorkerICE) connect(ctx context.Context, agent *icemaker.ThreadSafeAgent return } - w.log.Debugf("turn agent dial") - remoteConn, err := w.turnAgentDial(ctx, agent, remoteOfferAnswer) + w.log.Debugf("agent dial") + remoteConn, err := w.agentDial(ctx, agent, remoteOfferAnswer) if err != nil { w.log.Debugf("failed to dial the remote peer: %s", err) w.closeAgent(agent, w.agentDialerCancel) @@ -389,6 +389,17 @@ func (w *WorkerICE) injectPortForwardedCandidate(srflxCandidate ice.Candidate) { return } + // A forwarded candidate only makes sense for an IPv4 mapping, which + // translates a port on the gateway's address. An IPv6 pinhole translates + // nothing: it unblocks the address ICE already gathers as a host candidate, + // so there is no second address to advertise. Injecting one here would also + // paste an IPv6 address onto whichever server-reflexive candidate arrived + // first, which is usually IPv4. + if mapping.ExternalIP != nil && mapping.ExternalIP.To4() == nil { + w.log.Debugf("skipping port-forwarded candidate: %s mapping is IPv6-only", mapping.NATType) + return + } + w.muxAgent.Lock() if w.portForwardAttempted { w.muxAgent.Unlock() @@ -517,8 +528,8 @@ func (w *WorkerICE) onConnectionStateChange(agent *icemaker.ThreadSafeAgent, dia w.logSuccessfulPaths(agent) return case ice.ConnectionStateFailed, ice.ConnectionStateDisconnected, ice.ConnectionStateClosed: - // ice.ConnectionStateClosed happens when we recreate the agent. For the P2P to TURN switch important to - // notify the conn.onICEStateDisconnected changes to update the current used priority + // ice.ConnectionStateClosed happens when we recreate the agent. The P2P to relay switch requires + // notifying conn.onICEStateDisconnected so it can update the currently used priority. sessionChanged := w.closeAgent(agent, dialerCancel) @@ -532,7 +543,7 @@ func (w *WorkerICE) onConnectionStateChange(agent *icemaker.ThreadSafeAgent, dia } } -func (w *WorkerICE) turnAgentDial(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *OfferAnswer) (*ice.Conn, error) { +func (w *WorkerICE) agentDial(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *OfferAnswer) (*ice.Conn, error) { if isController(w.config) { return agent.Dial(ctx, remoteOfferAnswer.IceCredentials.UFrag, remoteOfferAnswer.IceCredentials.Pwd) } else { diff --git a/client/internal/portforward/manager.go b/client/internal/portforward/manager.go index b0680160c..7d5a4cb9e 100644 --- a/client/internal/portforward/manager.go +++ b/client/internal/portforward/manager.go @@ -10,10 +10,8 @@ import ( "sync" "time" - "github.com/libp2p/go-nat" + "github.com/netbirdio/go-nat" log "github.com/sirupsen/logrus" - - "github.com/netbirdio/netbird/client/internal/portforward/pcp" ) const ( @@ -168,6 +166,11 @@ func (m *Manager) setup(ctx context.Context) (nat.NAT, *Mapping, error) { if err != nil { return nil, nil, fmt.Errorf("create port mapping: %w", err) } + + // Only meaningful once a mapping has been attempted: that is what opens the + // pinhole and records its outcome. + logIPv6Pinhole(gateway) + return gateway, mapping, nil } @@ -265,7 +268,9 @@ func (m *Manager) checkHealthAndRecreate(ctx context.Context, gateway nat.NAT) b return false } - pcpNAT, ok := gateway.(*pcp.NAT) + // Assert on the interface, not on a concrete type: a dual-stack gateway is + // a wrapper around the IPv4 NAT, so a type assertion misses it. + checker, ok := gateway.(nat.HealthChecker) if !ok { return false } @@ -273,7 +278,7 @@ func (m *Manager) checkHealthAndRecreate(ctx context.Context, gateway nat.NAT) b ctx, cancel := context.WithTimeout(ctx, 10*time.Second) defer cancel() - epoch, serverRestarted, err := pcpNAT.CheckServerHealth(ctx) + epoch, serverRestarted, err := checker.CheckServerHealth(ctx) if err != nil { log.Debugf("PCP health check failed: %v", err) return false @@ -340,3 +345,18 @@ func (m *Manager) startTearDown(ctx context.Context) { func isPermanentLeaseRequired(err error) bool { return err != nil && upnpErrPermanentLeaseOnly.MatchString(err.Error()) } + +// logIPv6Pinhole reports the outcome of the IPv6 pinhole. Pinholes are best +// effort and never fail a mapping on their own, so this is the only way to see +// whether one was actually opened. +func logIPv6Pinhole(gateway nat.NAT) { + reporter, ok := gateway.(nat.IPv6PinholeReporter) + if !ok { + return + } + if err := reporter.IPv6PinholeError(); err != nil { + log.Warnf("IPv6 pinhole: %v", err) + return + } + log.Infof("IPv6 pinhole open") +} diff --git a/client/internal/portforward/pcp/client.go b/client/internal/portforward/pcp/client.go deleted file mode 100644 index f6d243ef9..000000000 --- a/client/internal/portforward/pcp/client.go +++ /dev/null @@ -1,408 +0,0 @@ -package pcp - -import ( - "context" - "crypto/rand" - "errors" - "fmt" - "net" - "net/netip" - "sync" - "time" - - log "github.com/sirupsen/logrus" -) - -const ( - defaultTimeout = 3 * time.Second - responseBufferSize = 128 - - // RFC 6887 Section 8.1.1 retry timing - initialRetryDelay = 3 * time.Second - maxRetryDelay = 1024 * time.Second - maxRetries = 4 // 3s + 6s + 12s + 24s = 45s total worst case -) - -// Client is a PCP protocol client. -// All methods are safe for concurrent use. -type Client struct { - gateway netip.Addr - timeout time.Duration - - mu sync.Mutex - // localIP caches the resolved local IP address. - localIP netip.Addr - // lastEpoch is the last observed server epoch value. - lastEpoch uint32 - // epochTime tracks when lastEpoch was received for state loss detection. - epochTime time.Time - // externalIP caches the external IP from the last successful MAP response. - externalIP netip.Addr - // epochStateLost is set when epoch indicates server restart. - epochStateLost bool -} - -// NewClient creates a new PCP client for the gateway at the given IP. -func NewClient(gateway net.IP) *Client { - addr, ok := netip.AddrFromSlice(gateway) - if !ok { - log.Debugf("invalid gateway IP: %v", gateway) - } - return &Client{ - gateway: addr.Unmap(), - timeout: defaultTimeout, - } -} - -// NewClientWithTimeout creates a new PCP client with a custom timeout. -func NewClientWithTimeout(gateway net.IP, timeout time.Duration) *Client { - addr, ok := netip.AddrFromSlice(gateway) - if !ok { - log.Debugf("invalid gateway IP: %v", gateway) - } - return &Client{ - gateway: addr.Unmap(), - timeout: timeout, - } -} - -// SetLocalIP sets the local IP address to use in PCP requests. -func (c *Client) SetLocalIP(ip net.IP) { - addr, ok := netip.AddrFromSlice(ip) - if !ok { - log.Debugf("invalid local IP: %v", ip) - } - c.mu.Lock() - c.localIP = addr.Unmap() - c.mu.Unlock() -} - -// Gateway returns the gateway IP address. -func (c *Client) Gateway() net.IP { - return c.gateway.AsSlice() -} - -// Announce sends a PCP ANNOUNCE request to discover PCP support. -// Returns the server's epoch time on success. -func (c *Client) Announce(ctx context.Context) (epoch uint32, err error) { - localIP, err := c.getLocalIP() - if err != nil { - return 0, fmt.Errorf("get local IP: %w", err) - } - - req := buildAnnounceRequest(localIP) - resp, err := c.sendRequest(ctx, req) - if err != nil { - return 0, fmt.Errorf("send announce: %w", err) - } - - parsed, err := parseResponse(resp) - if err != nil { - return 0, fmt.Errorf("parse announce response: %w", err) - } - - if parsed.ResultCode != ResultSuccess { - return 0, fmt.Errorf("PCP ANNOUNCE failed: %s", ResultCodeString(parsed.ResultCode)) - } - - c.mu.Lock() - if c.updateEpochLocked(parsed.Epoch) { - log.Warnf("PCP server epoch indicates state loss - mappings may need refresh") - } - c.mu.Unlock() - return parsed.Epoch, nil -} - -// AddPortMapping requests a port mapping from the PCP server. -func (c *Client) AddPortMapping(ctx context.Context, protocol string, internalPort int, lifetime time.Duration) (*MapResponse, error) { - return c.addPortMappingWithHint(ctx, protocol, internalPort, internalPort, netip.Addr{}, lifetime) -} - -// AddPortMappingWithHint requests a port mapping with suggested external port and IP. -// Use lifetime <= 0 to delete a mapping. -func (c *Client) AddPortMappingWithHint(ctx context.Context, protocol string, internalPort, suggestedExtPort int, suggestedExtIP net.IP, lifetime time.Duration) (*MapResponse, error) { - var extIP netip.Addr - if suggestedExtIP != nil { - var ok bool - extIP, ok = netip.AddrFromSlice(suggestedExtIP) - if !ok { - log.Debugf("invalid suggested external IP: %v", suggestedExtIP) - } - extIP = extIP.Unmap() - } - return c.addPortMappingWithHint(ctx, protocol, internalPort, suggestedExtPort, extIP, lifetime) -} - -func (c *Client) addPortMappingWithHint(ctx context.Context, protocol string, internalPort, suggestedExtPort int, suggestedExtIP netip.Addr, lifetime time.Duration) (*MapResponse, error) { - localIP, err := c.getLocalIP() - if err != nil { - return nil, fmt.Errorf("get local IP: %w", err) - } - - proto, err := protocolNumber(protocol) - if err != nil { - return nil, fmt.Errorf("parse protocol: %w", err) - } - - var nonce [12]byte - if _, err := rand.Read(nonce[:]); err != nil { - return nil, fmt.Errorf("generate nonce: %w", err) - } - - // Convert lifetime to seconds. Lifetime 0 means delete, so only apply - // default for positive durations that round to 0 seconds. - var lifetimeSec uint32 - if lifetime > 0 { - lifetimeSec = uint32(lifetime.Seconds()) - if lifetimeSec == 0 { - lifetimeSec = DefaultLifetime - } - } - - req := buildMapRequest(localIP, nonce, proto, uint16(internalPort), uint16(suggestedExtPort), suggestedExtIP, lifetimeSec) - - resp, err := c.sendRequest(ctx, req) - if err != nil { - return nil, fmt.Errorf("send map request: %w", err) - } - - mapResp, err := parseMapResponse(resp) - if err != nil { - return nil, fmt.Errorf("parse map response: %w", err) - } - - if mapResp.Nonce != nonce { - return nil, fmt.Errorf("nonce mismatch in response") - } - - if mapResp.Protocol != proto { - return nil, fmt.Errorf("protocol mismatch: requested %d, got %d", proto, mapResp.Protocol) - } - if mapResp.InternalPort != uint16(internalPort) { - return nil, fmt.Errorf("internal port mismatch: requested %d, got %d", internalPort, mapResp.InternalPort) - } - - if mapResp.ResultCode != ResultSuccess { - return nil, &Error{ - Code: mapResp.ResultCode, - Message: ResultCodeString(mapResp.ResultCode), - } - } - - c.mu.Lock() - if c.updateEpochLocked(mapResp.Epoch) { - log.Warnf("PCP server epoch indicates state loss - mappings may need refresh") - } - c.cacheExternalIPLocked(mapResp.ExternalIP) - c.mu.Unlock() - return mapResp, nil -} - -// DeletePortMapping removes a port mapping by requesting zero lifetime. -func (c *Client) DeletePortMapping(ctx context.Context, protocol string, internalPort int) error { - if _, err := c.addPortMappingWithHint(ctx, protocol, internalPort, 0, netip.Addr{}, 0); err != nil { - var pcpErr *Error - if errors.As(err, &pcpErr) && pcpErr.Code == ResultNotAuthorized { - return nil - } - return fmt.Errorf("delete mapping: %w", err) - } - return nil -} - -// GetExternalAddress returns the external IP address. -// First checks for a cached value from previous MAP responses. -// If not cached, creates a short-lived mapping to discover the external IP. -func (c *Client) GetExternalAddress(ctx context.Context) (net.IP, error) { - c.mu.Lock() - if c.externalIP.IsValid() { - ip := c.externalIP.AsSlice() - c.mu.Unlock() - return ip, nil - } - c.mu.Unlock() - - // Use an ephemeral port in the dynamic range (49152-65535). - // Port 0 is not valid with UDP/TCP protocols per RFC 6887. - ephemeralPort := 49152 + int(uint16(time.Now().UnixNano()))%(65535-49152) - - // Use minimal lifetime (1 second) for discovery. - resp, err := c.AddPortMapping(ctx, "udp", ephemeralPort, time.Second) - if err != nil { - return nil, fmt.Errorf("create temporary mapping: %w", err) - } - - if err := c.DeletePortMapping(ctx, "udp", ephemeralPort); err != nil { - log.Debugf("cleanup temporary PCP mapping: %v", err) - } - - return resp.ExternalIP.AsSlice(), nil -} - -// LastEpoch returns the last observed server epoch value. -// A decrease in epoch indicates the server may have restarted and mappings may be lost. -func (c *Client) LastEpoch() uint32 { - c.mu.Lock() - defer c.mu.Unlock() - return c.lastEpoch -} - -// EpochStateLost returns true if epoch state loss was detected and clears the flag. -func (c *Client) EpochStateLost() bool { - c.mu.Lock() - defer c.mu.Unlock() - lost := c.epochStateLost - c.epochStateLost = false - return lost -} - -// updateEpoch updates the epoch tracking and detects potential state loss. -// Returns true if state loss was detected (server likely restarted). -// Caller must hold c.mu. -func (c *Client) updateEpochLocked(newEpoch uint32) bool { - now := time.Now() - stateLost := false - - // RFC 6887 Section 8.5: Detect invalid epoch indicating server state loss. - // client_delta = time since last response - // server_delta = epoch change since last response - // Invalid if: client_delta+2 < server_delta - server_delta/16 - // OR: server_delta+2 < client_delta - client_delta/16 - // The +2 handles quantization, /16 (6.25%) handles clock drift. - if !c.epochTime.IsZero() && c.lastEpoch > 0 { - clientDelta := uint32(now.Sub(c.epochTime).Seconds()) - serverDelta := newEpoch - c.lastEpoch - - // Check for epoch going backwards or jumping unexpectedly. - // Subtraction is safe: serverDelta/16 is always <= serverDelta. - if clientDelta+2 < serverDelta-(serverDelta/16) || - serverDelta+2 < clientDelta-(clientDelta/16) { - stateLost = true - c.epochStateLost = true - } - } - - c.lastEpoch = newEpoch - c.epochTime = now - return stateLost -} - -// cacheExternalIP stores the external IP from a successful MAP response. -// Caller must hold c.mu. -func (c *Client) cacheExternalIPLocked(ip netip.Addr) { - if ip.IsValid() && !ip.IsUnspecified() { - c.externalIP = ip - } -} - -// sendRequest sends a PCP request with retries per RFC 6887 Section 8.1.1. -func (c *Client) sendRequest(ctx context.Context, req []byte) ([]byte, error) { - addr := &net.UDPAddr{IP: c.gateway.AsSlice(), Port: Port} - - var lastErr error - delay := initialRetryDelay - - for range maxRetries { - resp, err := c.sendOnce(ctx, addr, req) - if err == nil { - return resp, nil - } - lastErr = err - - if ctx.Err() != nil { - return nil, ctx.Err() - } - - // RFC 6887 Section 8.1.1: RT = (1 + RAND) * MIN(2 * RTprev, MRT) - // RAND is random between -0.1 and +0.1 - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-time.After(retryDelayWithJitter(delay)): - } - delay = min(delay*2, maxRetryDelay) - } - - return nil, fmt.Errorf("PCP request failed after %d retries: %w", maxRetries, lastErr) -} - -// retryDelayWithJitter applies RFC 6887 jitter: multiply by (1 + RAND) where RAND is [-0.1, +0.1]. -func retryDelayWithJitter(d time.Duration) time.Duration { - var b [1]byte - _, _ = rand.Read(b[:]) - // Convert byte to range [-0.1, +0.1]: (b/255 * 0.2) - 0.1 - jitter := (float64(b[0])/255.0)*0.2 - 0.1 - return time.Duration(float64(d) * (1 + jitter)) -} - -func (c *Client) sendOnce(ctx context.Context, addr *net.UDPAddr, req []byte) ([]byte, error) { - // Use ListenUDP instead of DialUDP to validate response source address per RFC 6887 §8.3. - conn, err := net.ListenUDP("udp", nil) - if err != nil { - return nil, fmt.Errorf("listen: %w", err) - } - defer func() { - if err := conn.Close(); err != nil { - log.Debugf("close UDP connection: %v", err) - } - }() - - timeout := c.timeout - if deadline, ok := ctx.Deadline(); ok { - if remaining := time.Until(deadline); remaining < timeout { - timeout = remaining - } - } - - if err := conn.SetDeadline(time.Now().Add(timeout)); err != nil { - return nil, fmt.Errorf("set deadline: %w", err) - } - - if _, err := conn.WriteToUDP(req, addr); err != nil { - return nil, fmt.Errorf("write: %w", err) - } - - resp := make([]byte, responseBufferSize) - n, from, err := conn.ReadFromUDP(resp) - if err != nil { - return nil, fmt.Errorf("read: %w", err) - } - - // RFC 6887 §8.3: Validate response came from expected PCP server. - if !from.IP.Equal(addr.IP) { - return nil, fmt.Errorf("response from unexpected source %s (expected %s)", from.IP, addr.IP) - } - - return resp[:n], nil -} - -func (c *Client) getLocalIP() (netip.Addr, error) { - c.mu.Lock() - defer c.mu.Unlock() - - if !c.localIP.IsValid() { - return netip.Addr{}, fmt.Errorf("local IP not set for gateway %s", c.gateway) - } - return c.localIP, nil -} - -func protocolNumber(protocol string) (uint8, error) { - switch protocol { - case "udp", "UDP": - return ProtoUDP, nil - case "tcp", "TCP": - return ProtoTCP, nil - default: - return 0, fmt.Errorf("unsupported protocol: %s", protocol) - } -} - -// Error represents a PCP error response. -type Error struct { - Code uint8 - Message string -} - -func (e *Error) Error() string { - return fmt.Sprintf("PCP error: %s (%d)", e.Message, e.Code) -} diff --git a/client/internal/portforward/pcp/client_test.go b/client/internal/portforward/pcp/client_test.go deleted file mode 100644 index 79f44a426..000000000 --- a/client/internal/portforward/pcp/client_test.go +++ /dev/null @@ -1,187 +0,0 @@ -package pcp - -import ( - "context" - "net" - "net/netip" - "testing" - "time" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestAddrConversion(t *testing.T) { - tests := []struct { - name string - addr netip.Addr - }{ - {"IPv4", netip.MustParseAddr("192.168.1.100")}, - {"IPv4 loopback", netip.MustParseAddr("127.0.0.1")}, - {"IPv6", netip.MustParseAddr("2001:db8::1")}, - {"IPv6 loopback", netip.MustParseAddr("::1")}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - b16 := addrTo16(tt.addr) - - recovered := addrFrom16(b16) - assert.Equal(t, tt.addr, recovered, "address should round-trip") - }) - } -} - -func TestBuildAnnounceRequest(t *testing.T) { - clientIP := netip.MustParseAddr("192.168.1.100") - req := buildAnnounceRequest(clientIP) - - require.Len(t, req, headerSize) - assert.Equal(t, byte(Version), req[0], "version") - assert.Equal(t, byte(OpAnnounce), req[1], "opcode") - - // Check client IP is properly encoded as IPv4-mapped IPv6 - assert.Equal(t, byte(0xff), req[18], "IPv4-mapped prefix byte 10") - assert.Equal(t, byte(0xff), req[19], "IPv4-mapped prefix byte 11") - assert.Equal(t, byte(192), req[20], "IP octet 1") - assert.Equal(t, byte(168), req[21], "IP octet 2") - assert.Equal(t, byte(1), req[22], "IP octet 3") - assert.Equal(t, byte(100), req[23], "IP octet 4") -} - -func TestBuildMapRequest(t *testing.T) { - clientIP := netip.MustParseAddr("192.168.1.100") - nonce := [12]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12} - req := buildMapRequest(clientIP, nonce, ProtoUDP, 51820, 51820, netip.Addr{}, 3600) - - require.Len(t, req, mapRequestSize) - assert.Equal(t, byte(Version), req[0], "version") - assert.Equal(t, byte(OpMap), req[1], "opcode") - - // Lifetime at bytes 4-7 - assert.Equal(t, uint32(3600), (uint32(req[4])<<24)|(uint32(req[5])<<16)|(uint32(req[6])<<8)|uint32(req[7]), "lifetime") - - // Nonce at bytes 24-35 - assert.Equal(t, nonce[:], req[24:36], "nonce") - - // Protocol at byte 36 - assert.Equal(t, byte(ProtoUDP), req[36], "protocol") - - // Internal port at bytes 40-41 - assert.Equal(t, uint16(51820), (uint16(req[40])<<8)|uint16(req[41]), "internal port") - - // External port at bytes 42-43 - assert.Equal(t, uint16(51820), (uint16(req[42])<<8)|uint16(req[43]), "external port") -} - -func TestParseResponse(t *testing.T) { - // Construct a valid ANNOUNCE response - resp := make([]byte, headerSize) - resp[0] = Version - resp[1] = OpAnnounce | OpReply - // Result code = 0 (success) - // Lifetime = 0 - // Epoch = 12345 - resp[8] = 0 - resp[9] = 0 - resp[10] = 0x30 - resp[11] = 0x39 - - parsed, err := parseResponse(resp) - require.NoError(t, err) - assert.Equal(t, uint8(Version), parsed.Version) - assert.Equal(t, uint8(OpAnnounce|OpReply), parsed.Opcode) - assert.Equal(t, uint8(ResultSuccess), parsed.ResultCode) - assert.Equal(t, uint32(12345), parsed.Epoch) -} - -func TestParseResponseErrors(t *testing.T) { - t.Run("too short", func(t *testing.T) { - _, err := parseResponse([]byte{1, 2, 3}) - assert.Error(t, err) - }) - - t.Run("wrong version", func(t *testing.T) { - resp := make([]byte, headerSize) - resp[0] = 1 // Wrong version - resp[1] = OpReply - _, err := parseResponse(resp) - assert.Error(t, err) - }) - - t.Run("missing reply bit", func(t *testing.T) { - resp := make([]byte, headerSize) - resp[0] = Version - resp[1] = OpAnnounce // Missing OpReply bit - _, err := parseResponse(resp) - assert.Error(t, err) - }) -} - -func TestResultCodeString(t *testing.T) { - assert.Equal(t, "SUCCESS", ResultCodeString(ResultSuccess)) - assert.Equal(t, "NOT_AUTHORIZED", ResultCodeString(ResultNotAuthorized)) - assert.Equal(t, "ADDRESS_MISMATCH", ResultCodeString(ResultAddressMismatch)) - assert.Contains(t, ResultCodeString(255), "UNKNOWN") -} - -func TestProtocolNumber(t *testing.T) { - proto, err := protocolNumber("udp") - require.NoError(t, err) - assert.Equal(t, uint8(ProtoUDP), proto) - - proto, err = protocolNumber("tcp") - require.NoError(t, err) - assert.Equal(t, uint8(ProtoTCP), proto) - - proto, err = protocolNumber("UDP") - require.NoError(t, err) - assert.Equal(t, uint8(ProtoUDP), proto) - - _, err = protocolNumber("icmp") - assert.Error(t, err) -} - -func TestClientCreation(t *testing.T) { - gateway := netip.MustParseAddr("192.168.1.1").AsSlice() - - client := NewClient(gateway) - assert.Equal(t, net.IP(gateway), client.Gateway()) - assert.Equal(t, defaultTimeout, client.timeout) - - clientWithTimeout := NewClientWithTimeout(gateway, 5*time.Second) - assert.Equal(t, 5*time.Second, clientWithTimeout.timeout) -} - -func TestNATType(t *testing.T) { - n := NewNAT(netip.MustParseAddr("192.168.1.1").AsSlice(), netip.MustParseAddr("192.168.1.100").AsSlice()) - assert.Equal(t, "PCP", n.Type()) -} - -// Integration test - skipped unless PCP_TEST_GATEWAY env is set -func TestClientIntegration(t *testing.T) { - t.Skip("Integration test - run manually with PCP_TEST_GATEWAY=") - - gateway := netip.MustParseAddr("10.0.1.1").AsSlice() // Change to your test gateway - localIP := netip.MustParseAddr("10.0.1.100").AsSlice() // Change to your local IP - - client := NewClient(gateway) - client.SetLocalIP(localIP) - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - // Test ANNOUNCE - epoch, err := client.Announce(ctx) - require.NoError(t, err) - t.Logf("Server epoch: %d", epoch) - - // Test MAP - resp, err := client.AddPortMapping(ctx, "udp", 51820, 1*time.Hour) - require.NoError(t, err) - t.Logf("Mapping: internal=%d external=%d externalIP=%s", - resp.InternalPort, resp.ExternalPort, resp.ExternalIP) - - // Cleanup - err = client.DeletePortMapping(ctx, "udp", 51820) - require.NoError(t, err) -} diff --git a/client/internal/portforward/pcp/nat.go b/client/internal/portforward/pcp/nat.go deleted file mode 100644 index 0e635b6c8..000000000 --- a/client/internal/portforward/pcp/nat.go +++ /dev/null @@ -1,222 +0,0 @@ -package pcp - -import ( - "context" - "fmt" - "net" - "net/netip" - "runtime" - "sync" - "time" - - log "github.com/sirupsen/logrus" - - "github.com/libp2p/go-nat" - "github.com/libp2p/go-netroute" -) - -var _ nat.NAT = (*NAT)(nil) - -// NAT implements the go-nat NAT interface using PCP. -// Supports dual-stack (IPv4 and IPv6) when available. -// All methods are safe for concurrent use. -// -// TODO: IPv6 pinholes use the local IPv6 address. If the address changes -// (e.g., due to SLAAC rotation or network change), the pinhole becomes stale -// and needs to be recreated with the new address. -type NAT struct { - client *Client - - mu sync.RWMutex - // client6 is the IPv6 PCP client, nil if IPv6 is unavailable. - client6 *Client - // localIP6 caches the local IPv6 address used for PCP requests. - localIP6 netip.Addr -} - -// NewNAT creates a new NAT instance backed by PCP. -func NewNAT(gateway, localIP net.IP) *NAT { - client := NewClient(gateway) - client.SetLocalIP(localIP) - return &NAT{ - client: client, - } -} - -// Type returns "PCP" as the NAT type. -func (n *NAT) Type() string { - return "PCP" -} - -// GetDeviceAddress returns the gateway IP address. -func (n *NAT) GetDeviceAddress() (net.IP, error) { - return n.client.Gateway(), nil -} - -// GetExternalAddress returns the external IP address. -func (n *NAT) GetExternalAddress() (net.IP, error) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - return n.client.GetExternalAddress(ctx) -} - -// GetInternalAddress returns the local IP address used to communicate with the gateway. -func (n *NAT) GetInternalAddress() (net.IP, error) { - addr, err := n.client.getLocalIP() - if err != nil { - return nil, err - } - return addr.AsSlice(), nil -} - -// AddPortMapping creates a port mapping on both IPv4 and IPv6 (if available). -func (n *NAT) AddPortMapping(ctx context.Context, protocol string, internalPort int, _ string, timeout time.Duration) (int, error) { - resp, err := n.client.AddPortMapping(ctx, protocol, internalPort, timeout) - if err != nil { - return 0, fmt.Errorf("add mapping: %w", err) - } - - n.mu.RLock() - client6 := n.client6 - localIP6 := n.localIP6 - n.mu.RUnlock() - - if client6 == nil { - return int(resp.ExternalPort), nil - } - - if _, err := client6.AddPortMapping(ctx, protocol, internalPort, timeout); err != nil { - log.Warnf("IPv6 PCP mapping failed (continuing with IPv4): %v", err) - return int(resp.ExternalPort), nil - } - - log.Infof("created IPv6 PCP pinhole: %s:%d", localIP6, internalPort) - return int(resp.ExternalPort), nil -} - -// DeletePortMapping removes a port mapping from both IPv4 and IPv6. -func (n *NAT) DeletePortMapping(ctx context.Context, protocol string, internalPort int) error { - err := n.client.DeletePortMapping(ctx, protocol, internalPort) - - n.mu.RLock() - client6 := n.client6 - n.mu.RUnlock() - - if client6 != nil { - if err6 := client6.DeletePortMapping(ctx, protocol, internalPort); err6 != nil { - log.Warnf("IPv6 PCP delete mapping failed: %v", err6) - } - } - - if err != nil { - return fmt.Errorf("delete mapping: %w", err) - } - return nil -} - -// CheckServerHealth sends an ANNOUNCE to verify the server is still responsive. -// Returns the current epoch and whether the server may have restarted (epoch state loss detected). -func (n *NAT) CheckServerHealth(ctx context.Context) (epoch uint32, serverRestarted bool, err error) { - epoch, err = n.client.Announce(ctx) - if err != nil { - return 0, false, fmt.Errorf("announce: %w", err) - } - return epoch, n.client.EpochStateLost(), nil -} - -// DiscoverPCP attempts to discover a PCP-capable gateway. -// Returns a NAT interface if PCP is supported, or an error otherwise. -// Discovers both IPv4 and IPv6 gateways when available. -func DiscoverPCP(ctx context.Context) (nat.NAT, error) { - gateway, localIP, err := getDefaultGateway() - if err != nil { - return nil, fmt.Errorf("get default gateway: %w", err) - } - - client := NewClient(gateway) - client.SetLocalIP(localIP) - if _, err := client.Announce(ctx); err != nil { - return nil, fmt.Errorf("PCP announce: %w", err) - } - - result := &NAT{client: client} - discoverIPv6(ctx, result) - - return result, nil -} - -func discoverIPv6(ctx context.Context, result *NAT) { - gateway6, localIP6, err := getDefaultGateway6() - if err != nil { - log.Debugf("IPv6 gateway discovery failed: %v", err) - return - } - - client6 := NewClient(gateway6) - client6.SetLocalIP(localIP6) - if _, err := client6.Announce(ctx); err != nil { - log.Debugf("PCP IPv6 announce failed: %v", err) - return - } - - addr, ok := netip.AddrFromSlice(localIP6) - if !ok { - log.Debugf("invalid IPv6 local IP: %v", localIP6) - return - } - result.mu.Lock() - result.client6 = client6 - result.localIP6 = addr - result.mu.Unlock() - log.Debugf("PCP IPv6 gateway discovered: %s (local: %s)", gateway6, localIP6) -} - -// getDefaultGateway returns the default IPv4 gateway and local IP using the system routing table. -func getDefaultGateway() (gateway net.IP, localIP net.IP, err error) { - router, err := netroute.New() - if err != nil { - return nil, nil, err - } - - dst := net.IPv4zero - if runtime.GOOS == "linux" || runtime.GOOS == "android" { - // go-netroute v0.4.0 rejects unspecified destinations client-side on Linux/Android. - // TODO: on android/ios, use platform APIs (ConnectivityManager.getLinkProperties / - // NWPathMonitor) when netlink-based lookup is restricted or unavailable. - dst = net.IPv4(0, 0, 0, 1) - } - _, gateway, localIP, err = router.Route(dst) - if err != nil { - return nil, nil, err - } - - if gateway == nil { - return nil, nil, nat.ErrNoNATFound - } - - return gateway, localIP, nil -} - -// getDefaultGateway6 returns the default IPv6 gateway IP address using the system routing table. -func getDefaultGateway6() (gateway net.IP, localIP net.IP, err error) { - router, err := netroute.New() - if err != nil { - return nil, nil, err - } - - dst := net.IPv6zero - if runtime.GOOS == "linux" || runtime.GOOS == "android" { - // ::2 - dst = net.IP{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2} - } - _, gateway, localIP, err = router.Route(dst) - if err != nil { - return nil, nil, err - } - - if gateway == nil { - return nil, nil, nat.ErrNoNATFound - } - - return gateway, localIP, nil -} diff --git a/client/internal/portforward/pcp/protocol.go b/client/internal/portforward/pcp/protocol.go deleted file mode 100644 index d81c50c8c..000000000 --- a/client/internal/portforward/pcp/protocol.go +++ /dev/null @@ -1,225 +0,0 @@ -// Package pcp implements the Port Control Protocol (RFC 6887). -// -// # Implemented Features -// -// - ANNOUNCE opcode: Discovers PCP server support -// - MAP opcode: Creates/deletes port mappings (IPv4 NAT) and firewall pinholes (IPv6) -// - Dual-stack: Simultaneous IPv4 and IPv6 support via separate clients -// - Nonce validation: Prevents response spoofing -// - Epoch tracking: Detects server restarts per Section 8.5 -// - RFC-compliant retry timing: 3s initial, exponential backoff to 1024s max (Section 8.1.1) -// -// # Not Implemented -// -// - PEER opcode: For outbound peer connections (not needed for inbound NAT traversal) -// - THIRD_PARTY option: For managing mappings on behalf of other devices -// - PREFER_FAILURE option: Requires exact external port or fail (IPv4 NAT only, not needed for IPv6 pinholing) -// - FILTER option: To restrict remote peer addresses -// -// These optional features are omitted because the primary use case is simple -// port forwarding for WireGuard, which only requires MAP with default behavior. -package pcp - -import ( - "encoding/binary" - "fmt" - "net/netip" -) - -const ( - // Version is the PCP protocol version (RFC 6887). - Version = 2 - - // Port is the standard PCP server port. - Port = 5351 - - // DefaultLifetime is the default requested mapping lifetime in seconds. - DefaultLifetime = 7200 // 2 hours - - // Header sizes - headerSize = 24 - mapPayloadSize = 36 - mapRequestSize = headerSize + mapPayloadSize // 60 bytes -) - -// Opcodes -const ( - OpAnnounce = 0 - OpMap = 1 - OpPeer = 2 - OpReply = 0x80 // OR'd with opcode in responses -) - -// Protocol numbers for MAP requests -const ( - ProtoUDP = 17 - ProtoTCP = 6 -) - -// Result codes (RFC 6887 Section 7.4) -const ( - ResultSuccess = 0 - ResultUnsuppVersion = 1 - ResultNotAuthorized = 2 - ResultMalformedRequest = 3 - ResultUnsuppOpcode = 4 - ResultUnsuppOption = 5 - ResultMalformedOption = 6 - ResultNetworkFailure = 7 - ResultNoResources = 8 - ResultUnsuppProtocol = 9 - ResultUserExQuota = 10 - ResultCannotProvideExt = 11 - ResultAddressMismatch = 12 - ResultExcessiveRemotePeers = 13 -) - -// ResultCodeString returns a human-readable string for a result code. -func ResultCodeString(code uint8) string { - switch code { - case ResultSuccess: - return "SUCCESS" - case ResultUnsuppVersion: - return "UNSUPP_VERSION" - case ResultNotAuthorized: - return "NOT_AUTHORIZED" - case ResultMalformedRequest: - return "MALFORMED_REQUEST" - case ResultUnsuppOpcode: - return "UNSUPP_OPCODE" - case ResultUnsuppOption: - return "UNSUPP_OPTION" - case ResultMalformedOption: - return "MALFORMED_OPTION" - case ResultNetworkFailure: - return "NETWORK_FAILURE" - case ResultNoResources: - return "NO_RESOURCES" - case ResultUnsuppProtocol: - return "UNSUPP_PROTOCOL" - case ResultUserExQuota: - return "USER_EX_QUOTA" - case ResultCannotProvideExt: - return "CANNOT_PROVIDE_EXTERNAL" - case ResultAddressMismatch: - return "ADDRESS_MISMATCH" - case ResultExcessiveRemotePeers: - return "EXCESSIVE_REMOTE_PEERS" - default: - return fmt.Sprintf("UNKNOWN(%d)", code) - } -} - -// Response represents a parsed PCP response header. -type Response struct { - Version uint8 - Opcode uint8 - ResultCode uint8 - Lifetime uint32 - Epoch uint32 -} - -// MapResponse contains the full response to a MAP request. -type MapResponse struct { - Response - Nonce [12]byte - Protocol uint8 - InternalPort uint16 - ExternalPort uint16 - ExternalIP netip.Addr -} - -// addrTo16 converts an address to its 16-byte IPv4-mapped IPv6 representation. -func addrTo16(addr netip.Addr) [16]byte { - if addr.Is4() { - return netip.AddrFrom4(addr.As4()).As16() - } - return addr.As16() -} - -// addrFrom16 extracts an address from a 16-byte representation, unmapping IPv4. -func addrFrom16(b [16]byte) netip.Addr { - return netip.AddrFrom16(b).Unmap() -} - -// buildAnnounceRequest creates a PCP ANNOUNCE request packet. -func buildAnnounceRequest(clientIP netip.Addr) []byte { - req := make([]byte, headerSize) - req[0] = Version - req[1] = OpAnnounce - mapped := addrTo16(clientIP) - copy(req[8:24], mapped[:]) - return req -} - -// buildMapRequest creates a PCP MAP request packet. -func buildMapRequest(clientIP netip.Addr, nonce [12]byte, protocol uint8, internalPort, suggestedExtPort uint16, suggestedExtIP netip.Addr, lifetime uint32) []byte { - req := make([]byte, mapRequestSize) - - // Header - req[0] = Version - req[1] = OpMap - binary.BigEndian.PutUint32(req[4:8], lifetime) - mapped := addrTo16(clientIP) - copy(req[8:24], mapped[:]) - - // MAP payload - copy(req[24:36], nonce[:]) - req[36] = protocol - binary.BigEndian.PutUint16(req[40:42], internalPort) - binary.BigEndian.PutUint16(req[42:44], suggestedExtPort) - if suggestedExtIP.IsValid() { - extMapped := addrTo16(suggestedExtIP) - copy(req[44:60], extMapped[:]) - } - - return req -} - -// parseResponse parses the common PCP response header. -func parseResponse(data []byte) (*Response, error) { - if len(data) < headerSize { - return nil, fmt.Errorf("response too short: %d bytes", len(data)) - } - - resp := &Response{ - Version: data[0], - Opcode: data[1], - ResultCode: data[3], // Byte 2 is reserved, byte 3 is result code (RFC 6887 §7.2) - Lifetime: binary.BigEndian.Uint32(data[4:8]), - Epoch: binary.BigEndian.Uint32(data[8:12]), - } - - if resp.Version != Version { - return nil, fmt.Errorf("unsupported PCP version: %d", resp.Version) - } - - if resp.Opcode&OpReply == 0 { - return nil, fmt.Errorf("response missing reply bit: opcode=0x%02x", resp.Opcode) - } - - return resp, nil -} - -// parseMapResponse parses a complete MAP response. -func parseMapResponse(data []byte) (*MapResponse, error) { - if len(data) < mapRequestSize { - return nil, fmt.Errorf("MAP response too short: %d bytes", len(data)) - } - - resp, err := parseResponse(data) - if err != nil { - return nil, fmt.Errorf("parse header: %w", err) - } - - mapResp := &MapResponse{ - Response: *resp, - Protocol: data[36], - InternalPort: binary.BigEndian.Uint16(data[40:42]), - ExternalPort: binary.BigEndian.Uint16(data[42:44]), - ExternalIP: addrFrom16([16]byte(data[44:60])), - } - copy(mapResp.Nonce[:], data[24:36]) - - return mapResp, nil -} diff --git a/client/internal/portforward/pinhole_test.go b/client/internal/portforward/pinhole_test.go new file mode 100644 index 000000000..46b07a9e7 --- /dev/null +++ b/client/internal/portforward/pinhole_test.go @@ -0,0 +1,116 @@ +//go:build !js + +package portforward + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/netbirdio/go-nat" + log "github.com/sirupsen/logrus" + "github.com/sirupsen/logrus/hooks/test" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// mockPinholeNAT is a gateway that also reports an IPv6 pinhole outcome, the +// shape a dual-stack gateway has. +type mockPinholeNAT struct { + *mockNAT + pinholeErr error +} + +func (m *mockPinholeNAT) IPv6PinholeError() error { + return m.pinholeErr +} + +func TestSetupLogsPinholeOutcome(t *testing.T) { + pinholeErr := errors.New("pcp ipv6: NOT_AUTHORIZED") + + tests := []struct { + name string + pinholeErr error + mappingErr error + wantLevel log.Level + wantText string + }{ + { + name: "an open pinhole is reported", + wantLevel: log.InfoLevel, + wantText: "IPv6 pinhole open", + }, + { + name: "a failed pinhole is reported without failing the mapping", + // The IPv4 mapping is what the caller asked for, so the pinhole + // failure surfaces only in the log. + pinholeErr: pinholeErr, + wantLevel: log.WarnLevel, + wantText: pinholeErr.Error(), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gateway := &mockPinholeNAT{mockNAT: newMockNAT(), pinholeErr: tt.pinholeErr} + hook := stubGatewayDiscovery(t, gateway) + + m := NewManager() + m.wgPort = 51820 + + _, mapping, err := m.setup(context.Background()) + + require.NoError(t, err) + require.NotNil(t, mapping) + + entry := findEntry(hook, tt.wantText) + require.NotNil(t, entry, "no log entry mentioning %q", tt.wantText) + assert.Equal(t, tt.wantLevel, entry.Level) + }) + } + + t.Run("a failed mapping reports no pinhole outcome", func(t *testing.T) { + // Nothing opened the pinhole, so whatever it currently reports says + // nothing about this attempt. + gateway := &mockPinholeNAT{mockNAT: newMockNAT()} + gateway.addMappingErr = errors.New("gateway refused") + hook := stubGatewayDiscovery(t, gateway) + + m := NewManager() + m.wgPort = 51820 + + _, _, err := m.setup(context.Background()) + + require.Error(t, err) + assert.Nil(t, findEntry(hook, "IPv6 pinhole")) + }) +} + +// stubGatewayDiscovery makes discovery return gateway and captures log output. +func stubGatewayDiscovery(t *testing.T, gateway nat.NAT) *test.Hook { + t.Helper() + + orig := discoverGateway + discoverGateway = func(context.Context) (nat.NAT, error) { return gateway, nil } + t.Cleanup(func() { discoverGateway = orig }) + + hook := test.NewGlobal() + origLevel := log.GetLevel() + log.SetLevel(log.DebugLevel) + t.Cleanup(func() { + hook.Reset() + log.SetLevel(origLevel) + }) + + return hook +} + +func findEntry(hook *test.Hook, substr string) *log.Entry { + for _, entry := range hook.AllEntries() { + if strings.Contains(entry.Message, substr) { + return entry + } + } + return nil +} diff --git a/client/internal/portforward/state.go b/client/internal/portforward/state.go index b1315cdc0..a21368e58 100644 --- a/client/internal/portforward/state.go +++ b/client/internal/portforward/state.go @@ -4,27 +4,94 @@ package portforward import ( "context" + "errors" "fmt" + "time" - "github.com/libp2p/go-nat" + "github.com/netbirdio/go-nat" + "github.com/netbirdio/go-nat/pcp" log "github.com/sirupsen/logrus" - - "github.com/netbirdio/netbird/client/internal/portforward/pcp" ) // discoverGateway is the function used for NAT gateway discovery. // It can be replaced in tests to avoid real network operations. -// Tries PCP first, then falls back to NAT-PMP/UPnP. var discoverGateway = defaultDiscoverGateway -func defaultDiscoverGateway(ctx context.Context) (nat.NAT, error) { - pcpGateway, err := pcp.DiscoverPCP(ctx) - if err == nil { - return pcpGateway, nil - } - log.Debugf("PCP discovery failed: %v, trying NAT-PMP/UPnP", err) +// pinholeDiscoveryTimeout is the slice of the discovery budget held back for +// the IPv6 pinhole probe. +// +// Sizing it is coarser than it looks: PCP retransmits on a 3s socket timeout +// and a 3s first backoff, so a second attempt needs about 9s. Anything from +// roughly 1s to 8s therefore buys exactly one attempt, and this only sets how +// long that attempt waits. A PCP server sits on the local link and answers in +// milliseconds, so 3s is margin rather than need, and the rest is left to +// gateway discovery, whose multicast SSDP search alone takes 5s. A probe lost +// to a dropped packet is retried by the next discovery round. +// +// It is a variable so tests can shorten it. +var pinholeDiscoveryTimeout = 3 * time.Second - return nat.DiscoverGateway(ctx) +// Discovery entry points, as variables so tests can drive the fallback without +// touching the network. +var ( + discoverNATGateway = nat.DiscoverGateway + + discoverPCPPinhole = func(ctx context.Context) (nat.NAT, error) { + pinhole, err := pcp.DiscoverPCP(ctx) + if err != nil { + return nil, err + } + return pinhole, nil + } +) + +// defaultDiscoverGateway finds a gateway that can make the WireGuard port +// reachable. DiscoverGateway prefers PCP for IPv4, races UPnP and NAT-PMP +// behind it, and attaches an IPv6 pinhole independently of which IPv4 protocol +// wins. +// +// It reports no gateway on a network offering only IPv6, having no IPv4 mapping +// to attach a pinhole to. Such a network still needs one: there is no +// translation to traverse, but the router drops inbound IPv6 until something +// opens it. Fall back to PCP alone, which yields a gateway holding just the +// pinhole. +func defaultDiscoverGateway(ctx context.Context) (nat.NAT, error) { + gatewayCtx, cancel := reserveForPinhole(ctx) + defer cancel() + + gateway, err := discoverNATGateway(gatewayCtx) + if err == nil { + return gateway, nil + } + if !errors.Is(err, nat.ErrNoNATFound) { + return nil, err + } + + pinhole, pinholeErr := discoverPCPPinhole(ctx) + if pinholeErr != nil { + log.Debugf("no IPv6 pinhole after %v: %v", err, pinholeErr) + return nil, err + } + + log.Infof("no IPv4 gateway, continuing with an IPv6 pinhole only") + return pinhole, nil +} + +// reserveForPinhole shortens ctx so that a pinhole probe still has time to run +// afterwards. Finding nothing takes gateway discovery everything it is given, +// so on the unshortened context the probe would start already expired. A budget +// too small to divide is left to gateway discovery, which is the likelier win. +func reserveForPinhole(ctx context.Context) (context.Context, context.CancelFunc) { + deadline, ok := ctx.Deadline() + if !ok { + return context.WithCancel(ctx) + } + + remaining := time.Until(deadline) + if remaining <= pinholeDiscoveryTimeout { + return context.WithCancel(ctx) + } + return context.WithTimeout(ctx, remaining-pinholeDiscoveryTimeout) } // State is persisted only for crash recovery cleanup diff --git a/client/internal/portforward/state_test.go b/client/internal/portforward/state_test.go new file mode 100644 index 000000000..8a584eecb --- /dev/null +++ b/client/internal/portforward/state_test.go @@ -0,0 +1,140 @@ +//go:build !js + +package portforward + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/netbirdio/go-nat" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// stubDiscovery replaces both discovery entry points for the duration of a +// test. gatewayDelay simulates gateway discovery spending everything it is +// given before reporting that it found nothing. +func stubDiscovery(t *testing.T, gateway nat.NAT, gatewayErr error, gatewayDelay time.Duration, pinhole nat.NAT, pinholeErr error) { + t.Helper() + + origGateway, origPinhole := discoverNATGateway, discoverPCPPinhole + discoverNATGateway = func(ctx context.Context) (nat.NAT, error) { + if gatewayDelay > 0 { + select { + case <-time.After(gatewayDelay): + case <-ctx.Done(): + } + } + return gateway, gatewayErr + } + discoverPCPPinhole = func(ctx context.Context) (nat.NAT, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + return pinhole, pinholeErr + } + + t.Cleanup(func() { discoverNATGateway, discoverPCPPinhole = origGateway, origPinhole }) +} + +func TestDefaultDiscoverGateway(t *testing.T) { + ipv4Gateway := &mockNAT{natType: "PCP+PCPv6"} + ipv6Pinhole := &mockNAT{natType: "PCP"} + otherErr := errors.New("routing table unavailable") + + t.Run("an IPv4 gateway is used as is", func(t *testing.T) { + stubDiscovery(t, ipv4Gateway, nil, 0, ipv6Pinhole, nil) + + got, err := defaultDiscoverGateway(context.Background()) + + require.NoError(t, err) + assert.Same(t, ipv4Gateway, got) + }) + + t.Run("no IPv4 gateway still opens an IPv6 pinhole", func(t *testing.T) { + stubDiscovery(t, nil, nat.ErrNoNATFound, 0, ipv6Pinhole, nil) + + got, err := defaultDiscoverGateway(context.Background()) + + require.NoError(t, err) + assert.Same(t, ipv6Pinhole, got) + }) + + t.Run("no gateway and no pinhole reports the original failure", func(t *testing.T) { + stubDiscovery(t, nil, nat.ErrNoNATFound, 0, nil, errors.New("no IPv6 route")) + + got, err := defaultDiscoverGateway(context.Background()) + + assert.Nil(t, got) + assert.ErrorIs(t, err, nat.ErrNoNATFound, "the pinhole failure must not mask why no gateway was found") + }) + + t.Run("a failure other than no-gateway is reported as is", func(t *testing.T) { + stubDiscovery(t, nil, otherErr, 0, ipv6Pinhole, nil) + + got, err := defaultDiscoverGateway(context.Background()) + + assert.Nil(t, got) + assert.ErrorIs(t, err, otherErr) + }) + + t.Run("the pinhole survives gateway discovery using its whole budget", func(t *testing.T) { + // On one shared context the probe would start already expired, which is + // how this failed against a real gateway. + reserve := 50 * time.Millisecond + origReserve := pinholeDiscoveryTimeout + pinholeDiscoveryTimeout = reserve + t.Cleanup(func() { pinholeDiscoveryTimeout = origReserve }) + + budget := 4 * reserve + ctx, cancel := context.WithTimeout(context.Background(), budget) + defer cancel() + + stubDiscovery(t, nil, nat.ErrNoNATFound, budget, ipv6Pinhole, nil) + + got, err := defaultDiscoverGateway(ctx) + + require.NoError(t, err) + assert.Same(t, ipv6Pinhole, got) + }) +} + +func TestReserveForPinhole(t *testing.T) { + origReserve := pinholeDiscoveryTimeout + pinholeDiscoveryTimeout = time.Second + t.Cleanup(func() { pinholeDiscoveryTimeout = origReserve }) + + t.Run("a budget is divided", func(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + gatewayCtx, cancelGateway := reserveForPinhole(ctx) + defer cancelGateway() + + deadline, ok := gatewayCtx.Deadline() + require.True(t, ok) + assert.InDelta(t, 9*time.Second, time.Until(deadline), float64(500*time.Millisecond)) + }) + + t.Run("a budget too small to divide is left whole", func(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) + defer cancel() + + gatewayCtx, cancelGateway := reserveForPinhole(ctx) + defer cancelGateway() + + deadline, ok := gatewayCtx.Deadline() + require.True(t, ok) + assert.InDelta(t, 500*time.Millisecond, time.Until(deadline), float64(100*time.Millisecond)) + }) + + t.Run("no deadline stays unbounded", func(t *testing.T) { + gatewayCtx, cancelGateway := reserveForPinhole(context.Background()) + defer cancelGateway() + + _, ok := gatewayCtx.Deadline() + assert.False(t, ok) + }) +} diff --git a/client/internal/profilemanager/config.go b/client/internal/profilemanager/config.go index ed2f21999..e1668238e 100644 --- a/client/internal/profilemanager/config.go +++ b/client/internal/profilemanager/config.go @@ -96,6 +96,7 @@ type ConfigInput struct { BlockLANAccess *bool BlockInbound *bool DisableIPv6 *bool + SyncMessageVersion *int DisableNotifications *bool @@ -137,6 +138,7 @@ type Config struct { BlockLANAccess bool BlockInbound bool DisableIPv6 bool + SyncMessageVersion *int DisableNotifications *bool @@ -587,6 +589,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } + if input.SyncMessageVersion != nil && *input.SyncMessageVersion != *config.SyncMessageVersion { + log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion) + *config.SyncMessageVersion = *input.SyncMessageVersion + updated = true + } + if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) { if *input.DisableNotifications { log.Infof("disabling notifications") @@ -738,6 +746,13 @@ func (config *Config) applyMDMPolicy(policy *mdm.Policy) { // appended for https or ":80" for http. The serviceName parameter is // used to contextualise error messages. On success returns the parsed // *url.URL; on failure returns a non-nil error. +// ParseServiceURL normalises a service URL exactly as the config layer does when +// it stores one, so callers comparing a requested URL against a stored one do not +// have to reimplement the scheme validation and default-port handling. +func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) { + return parseURL(serviceName, serviceURL) +} + func parseURL(serviceName, serviceURL string) (*url.URL, error) { parsedMgmtURL, err := url.ParseRequestURI(serviceURL) if err != nil { diff --git a/client/internal/profilemanager/prefs.go b/client/internal/profilemanager/prefs.go new file mode 100644 index 000000000..5613b0be3 --- /dev/null +++ b/client/internal/profilemanager/prefs.go @@ -0,0 +1,130 @@ +package profilemanager + +import ( + "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sync" + + "github.com/netbirdio/netbird/util" +) + +const prefsFileSuffix = ".prefs.json" + +var prefsMu sync.Mutex + +// Prefs is a namespaced per-profile preference store backed by a single JSON +// file next to the profile config; it is deleted together with the profile. +type Prefs struct { + path string +} + +// ProfilePrefs returns the preference store of the profile identified by id. +func (s *ServiceManager) ProfilePrefs(id ID, username string) (*Prefs, error) { + if !IsValidProfileFilenameStem(id) { + return nil, fmt.Errorf("invalid profile ID: %q", id) + } + if id == defaultProfileName { + return &Prefs{path: filepath.Join(filepath.Dir(DefaultConfigPath), id.String()+prefsFileSuffix)}, nil + } + configDir, err := s.getConfigDir(username) + if err != nil { + return nil, fmt.Errorf("get config directory for user %s: %w", username, err) + } + return &Prefs{path: filepath.Join(configDir, id.String()+prefsFileSuffix)}, nil +} + +// Get unmarshals the namespace section into v and reports whether it exists. +func (p *Prefs) Get(namespace string, v any) (bool, error) { + if namespace == "" { + return false, fmt.Errorf("empty prefs namespace") + } + + prefsMu.Lock() + defer prefsMu.Unlock() + + sections, err := readPrefsFile(p.path) + if err != nil { + return false, err + } + raw, ok := sections[namespace] + if !ok { + return false, nil + } + if err := json.Unmarshal(raw, v); err != nil { + return false, fmt.Errorf("decode prefs namespace %q: %w", namespace, err) + } + return true, nil +} + +// Put stores v as the namespace section, replacing any previous value. +func (p *Prefs) Put(namespace string, v any) error { + if namespace == "" { + return fmt.Errorf("empty prefs namespace") + } + raw, err := json.Marshal(v) + if err != nil { + return fmt.Errorf("encode prefs namespace %q: %w", namespace, err) + } + + prefsMu.Lock() + defer prefsMu.Unlock() + + sections, err := readPrefsFile(p.path) + if err != nil { + return err + } + sections[namespace] = raw + return writePrefsFile(p.path, sections) +} + +// Remove deletes the namespace section; a missing one is not an error. +func (p *Prefs) Remove(namespace string) error { + if namespace == "" { + return fmt.Errorf("empty prefs namespace") + } + + prefsMu.Lock() + defer prefsMu.Unlock() + + sections, err := readPrefsFile(p.path) + if err != nil { + return err + } + if _, ok := sections[namespace]; !ok { + return nil + } + delete(sections, namespace) + return writePrefsFile(p.path, sections) +} + +func removePrefsFile(path string) error { + prefsMu.Lock() + defer prefsMu.Unlock() + return os.Remove(path) +} + +func readPrefsFile(path string) (map[string]json.RawMessage, error) { + data, err := os.ReadFile(path) + if os.IsNotExist(err) { + return map[string]json.RawMessage{}, nil + } + if err != nil { + return nil, fmt.Errorf("read prefs: %w", err) + } + + sections := map[string]json.RawMessage{} + if err := json.Unmarshal(data, §ions); err != nil { + return nil, fmt.Errorf("decode prefs: %w", err) + } + return sections, nil +} + +func writePrefsFile(path string, sections map[string]json.RawMessage) error { + if err := util.WriteJsonWithRestrictedPermission(context.Background(), path, sections); err != nil { + return fmt.Errorf("write prefs: %w", err) + } + return nil +} diff --git a/client/internal/profilemanager/prefs_test.go b/client/internal/profilemanager/prefs_test.go new file mode 100644 index 000000000..692ade70f --- /dev/null +++ b/client/internal/profilemanager/prefs_test.go @@ -0,0 +1,138 @@ +package profilemanager + +import ( + "errors" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type testPrefsSection struct { + Mode uint8 `json:"mode"` + Dest string `json:"dest"` +} + +func TestProfilePrefs_RoundTrip(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + prefs, err := sm.ProfilePrefs(created.ID, username) + require.NoError(t, err) + + require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 2, Dest: "/tmp/x"})) + require.NoError(t, prefs.Put("other", map[string]int{"n": 1})) + + var got testPrefsSection + found, err := prefs.Get("filedrop", &got) + require.NoError(t, err) + assert.True(t, found) + assert.Equal(t, testPrefsSection{Mode: 2, Dest: "/tmp/x"}, got) + + var other map[string]int + found, err = prefs.Get("other", &other) + require.NoError(t, err) + assert.True(t, found) + assert.Equal(t, map[string]int{"n": 1}, other) + }) +} + +func TestProfilePrefs_GetMissingNamespace(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + prefs, err := sm.ProfilePrefs(created.ID, username) + require.NoError(t, err) + + var got testPrefsSection + found, err := prefs.Get("filedrop", &got) + require.NoError(t, err) + assert.False(t, found) + }) +} + +func TestProfilePrefs_RemoveNamespace(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + prefs, err := sm.ProfilePrefs(created.ID, username) + require.NoError(t, err) + + require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 1})) + require.NoError(t, prefs.Put("other", map[string]int{"n": 1})) + require.NoError(t, prefs.Remove("filedrop")) + require.NoError(t, prefs.Remove("missing")) + + var got testPrefsSection + found, err := prefs.Get("filedrop", &got) + require.NoError(t, err) + assert.False(t, found) + + var other map[string]int + found, err = prefs.Get("other", &other) + require.NoError(t, err) + assert.True(t, found) + assert.Equal(t, map[string]int{"n": 1}, other) + }) +} + +func TestProfilePrefs_RejectsInvalidID(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + _, err := sm.ProfilePrefs("../escape", username) + assert.Error(t, err) + }) +} + +func TestProfilePrefs_RejectsEmptyNamespace(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + prefs, err := sm.ProfilePrefs(created.ID, username) + require.NoError(t, err) + + _, err = prefs.Get("", &testPrefsSection{}) + assert.Error(t, err) + assert.Error(t, prefs.Put("", testPrefsSection{})) + assert.Error(t, prefs.Remove("")) + }) +} + +func TestProfilePrefs_DefaultProfile(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + prefs, err := sm.ProfilePrefs(defaultProfileName, username) + require.NoError(t, err) + + require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 1})) + + expected := filepath.Join(filepath.Dir(DefaultConfigPath), "default"+prefsFileSuffix) + _, err = os.Stat(expected) + require.NoError(t, err) + }) +} + +func TestRemoveProfile_DeletesPrefsFile(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + prefs, err := sm.ProfilePrefs(created.ID, username) + require.NoError(t, err) + require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 2})) + + configDir, err := sm.getConfigDir(username) + require.NoError(t, err) + prefsPath := filepath.Join(configDir, created.ID.String()+prefsFileSuffix) + _, err = os.Stat(prefsPath) + require.NoError(t, err) + + require.NoError(t, sm.RemoveProfile(created.ID, username)) + _, err = os.Stat(prefsPath) + assert.True(t, errors.Is(err, os.ErrNotExist), "prefs file should be removed") + }) +} diff --git a/client/internal/profilemanager/service.go b/client/internal/profilemanager/service.go index 696a60310..ec287f01a 100644 --- a/client/internal/profilemanager/service.go +++ b/client/internal/profilemanager/service.go @@ -420,6 +420,11 @@ func (s *ServiceManager) RemoveProfile(id ID, username string) error { log.Warnf("failed to remove profile state file %s: %v", stateFile, err) } + prefsFile := filepath.Join(filepath.Dir(target.Path), id.String()+prefsFileSuffix) + if err := removePrefsFile(prefsFile); err != nil && !os.IsNotExist(err) { + log.Warnf("failed to remove profile prefs file %s: %v", prefsFile, err) + } + return nil } diff --git a/client/internal/profilemanager/state.go b/client/internal/profilemanager/state.go index 9e9577796..ddb5dd056 100644 --- a/client/internal/profilemanager/state.go +++ b/client/internal/profilemanager/state.go @@ -45,12 +45,35 @@ func (pm *ProfileManager) GetProfileState(id ID) (*ProfileState, error) { return &state, nil } -func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { +// SetProfileState writes the state file of the profile identified by id. Prefer +// it over SetActiveProfileState whenever the caller knows which profile the data +// belongs to: an SSO login spans seconds of user interaction, and the active +// profile can change during it, which would file the account email under +// whichever profile happened to be active when the flow returned. +func (pm *ProfileManager) SetProfileState(id ID, state *ProfileState) error { configDir, err := getConfigDir() if err != nil { return fmt.Errorf("get config directory: %w", err) } + if id == "" { + return fmt.Errorf("empty profile ID") + } + if id != defaultProfileName && !IsValidProfileFilenameStem(id) { + return fmt.Errorf("invalid profile ID: %q", id) + } + + stateFile := filepath.Join(configDir, id.String()+".state.json") + if err := util.WriteJsonWithRestrictedPermission(context.Background(), stateFile, state); err != nil { + return fmt.Errorf("write profile state: %w", err) + } + + return nil +} + +// SetActiveProfileState writes the state file of whichever profile is active at +// call time. Use SetProfileState when the target profile is known. +func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { activeProf, err := pm.GetActiveProfile() if err != nil { if errors.Is(err, ErrNoActiveProfile) { @@ -59,25 +82,15 @@ func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { return fmt.Errorf("get active profile: %w", err) } - id := activeProf.ID - if id != defaultProfileName && !IsValidProfileFilenameStem(id) { - return fmt.Errorf("invalid active profile ID: %q", id) - } - - stateFile := filepath.Join(configDir, id.String()+".state.json") - err = util.WriteJsonWithRestrictedPermission(context.Background(), stateFile, state) - if err != nil { - return fmt.Errorf("write profile state: %w", err) - } - - return nil + return pm.SetProfileState(activeProf.ID, state) } // RemoveProfileState deletes the per-profile state file (which holds the // account email used for the SSO login hint and the UI display). Called after -// a successful logout so a logged-out profile no longer shows a stale account -// email. The state file only stores the email, so deleting it is equivalent to -// clearing it; the next SSO login recreates it. A missing file is not an error. +// profile removal; logout keeps the file so the next login can pass the email +// as the login_hint. The state file only stores the email, so deleting it is +// equivalent to clearing it; the next SSO login recreates it. A missing file +// is not an error. func (pm *ProfileManager) RemoveProfileState(profileName string) error { configDir, err := getConfigDir() if err != nil { diff --git a/client/internal/routemanager/dnsinterceptor/handler.go b/client/internal/routemanager/dnsinterceptor/handler.go index b784cc274..d20f4b944 100644 --- a/client/internal/routemanager/dnsinterceptor/handler.go +++ b/client/internal/routemanager/dnsinterceptor/handler.go @@ -95,7 +95,7 @@ func (d *DnsInterceptor) RemoveRoute() error { // AllowedIPs should use real IPs if d.currentPeerKey != "" { - if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err)) } } @@ -172,7 +172,7 @@ func (d *DnsInterceptor) removeAllowedIP(realPrefix netip.Prefix) error { } // AllowedIPs use real IPs - if _, err := d.allowedIPsRefcounter.Decrement(realPrefix); err != nil { + if _, err := d.allowedIPsRefcounter.Decrement(realPrefix, d.currentPeerKey); err != nil { return fmt.Errorf("remove allowed IP %s: %v", realPrefix, err) } @@ -205,7 +205,7 @@ func (d *DnsInterceptor) RemoveAllowedIPs() error { for _, prefixes := range d.interceptedDomains { for _, prefix := range prefixes { // AllowedIPs use real IPs - if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err)) } } @@ -479,7 +479,7 @@ func (d *DnsInterceptor) removeDNATMappings(realPrefixes []netip.Prefix, logger // internalDnatFw checks if the firewall supports internal DNAT func (d *DnsInterceptor) internalDnatFw() (internalDNATer, bool) { - if d.firewall == nil || runtime.GOOS != "android" { + if d.firewall == nil || d.fakeIPManager == nil || runtime.GOOS != "android" { return nil, false } fw, ok := d.firewall.(internalDNATer) diff --git a/client/internal/routemanager/dynamic/route.go b/client/internal/routemanager/dynamic/route.go index f0efd7b22..bb3b1c59c 100644 --- a/client/internal/routemanager/dynamic/route.go +++ b/client/internal/routemanager/dynamic/route.go @@ -135,7 +135,7 @@ func (r *Route) RemoveAllowedIPs() error { var merr *multierror.Error for _, domainPrefixes := range r.dynamicDomains { for _, prefix := range domainPrefixes { - if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err)) } } @@ -185,7 +185,7 @@ func (r *Route) startResolver(ctx context.Context) { } func (r *Route) update(ctx context.Context) error { - resolved, err := r.resolveDomains() + resolved, err := r.resolveDomains(ctx) if err != nil { if len(resolved) == 0 { return fmt.Errorf("resolve domains: %w", err) @@ -199,9 +199,9 @@ func (r *Route) update(ctx context.Context) error { return nil } -func (r *Route) resolveDomains() (domainMap, error) { +func (r *Route) resolveDomains(ctx context.Context) (domainMap, error) { results := make(chan resolveResult) - go r.resolve(results) + go r.resolve(ctx, results) resolved := domainMap{} var merr *multierror.Error @@ -217,7 +217,7 @@ func (r *Route) resolveDomains() (domainMap, error) { return resolved, nberrors.FormatErrorOrNil(merr) } -func (r *Route) resolve(results chan resolveResult) { +func (r *Route) resolve(ctx context.Context, results chan resolveResult) { var wg sync.WaitGroup for _, d := range r.route.Domains { @@ -225,10 +225,10 @@ func (r *Route) resolve(results chan resolveResult) { go func(domain domain.Domain) { defer wg.Done() - ips, err := r.getIPsFromResolver(domain) + ips, err := r.getIPsFromResolver(ctx, domain) if err != nil { log.Tracef("Failed to resolve domain %s with private resolver: %v", domain.SafeString(), err) - ips, err = net.LookupIP(domain.PunycodeString()) + ips, err = lookupHostIPs(ctx, domain) if err != nil { results <- resolveResult{domain: domain, err: fmt.Errorf("resolve d %s: %w", domain.SafeString(), err)} return @@ -320,7 +320,7 @@ func (r *Route) removeRoutes(prefixes []netip.Prefix) ([]netip.Prefix, error) { merr = multierror.Append(merr, fmt.Errorf("remove dynamic route for IP %s: %w", prefix, err)) } if r.currentPeerKey != "" { - if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err)) } } @@ -364,6 +364,20 @@ func determinePrefixChanges(oldPrefixes, newPrefixes []netip.Prefix) (toAdd, toR return } +// lookupHostIPs resolves d via the system resolver, honoring ctx cancellation. +func lookupHostIPs(ctx context.Context, d domain.Domain) ([]net.IP, error) { + addrs, err := net.DefaultResolver.LookupIPAddr(ctx, d.PunycodeString()) + if err != nil { + return nil, err + } + + ips := make([]net.IP, 0, len(addrs)) + for _, addr := range addrs { + ips = append(ips, addr.IP) + } + return ips, nil +} + func combinePrefixes(oldPrefixes, removedPrefixes, addedPrefixes []netip.Prefix) []netip.Prefix { prefixSet := make(map[netip.Prefix]struct{}) for _, prefix := range oldPrefixes { diff --git a/client/internal/routemanager/dynamic/route_generic.go b/client/internal/routemanager/dynamic/route_generic.go index 56fd63fba..8bc2dd3df 100644 --- a/client/internal/routemanager/dynamic/route_generic.go +++ b/client/internal/routemanager/dynamic/route_generic.go @@ -3,11 +3,12 @@ package dynamic import ( + "context" "net" "github.com/netbirdio/netbird/shared/management/domain" ) -func (r *Route) getIPsFromResolver(domain domain.Domain) ([]net.IP, error) { - return net.LookupIP(domain.PunycodeString()) +func (r *Route) getIPsFromResolver(ctx context.Context, domain domain.Domain) ([]net.IP, error) { + return lookupHostIPs(ctx, domain) } diff --git a/client/internal/routemanager/dynamic/route_ios.go b/client/internal/routemanager/dynamic/route_ios.go index 1ae281d56..6a3d262b8 100644 --- a/client/internal/routemanager/dynamic/route_ios.go +++ b/client/internal/routemanager/dynamic/route_ios.go @@ -3,6 +3,7 @@ package dynamic import ( + "context" "fmt" "net" "time" @@ -16,7 +17,7 @@ import ( const dialTimeout = 10 * time.Second -func (r *Route) getIPsFromResolver(domain domain.Domain) ([]net.IP, error) { +func (r *Route) getIPsFromResolver(ctx context.Context, domain domain.Domain) ([]net.IP, error) { privateClient, err := nbdns.GetClientPrivate(r.wgInterface, r.resolverAddr.Addr(), dialTimeout) if err != nil { return nil, fmt.Errorf("error while creating private client: %s", err) @@ -32,7 +33,7 @@ func (r *Route) getIPsFromResolver(domain domain.Domain) ([]net.IP, error) { msg := new(dns.Msg) msg.SetQuestion(fqdn, qtype) - response, _, err := nbdns.ExchangeWithFallback(nil, privateClient, msg, r.resolverAddr.String()) + response, _, err := nbdns.ExchangeWithFallback(ctx, privateClient, msg, r.resolverAddr.String()) if err != nil { if queryErr == nil { queryErr = fmt.Errorf("DNS query for %s (type %d) after %s: %w", domain.SafeString(), qtype, time.Since(startTime), err) diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate.go b/client/internal/routemanager/ipfwdstate/ipfwdstate.go index 2be1c2ae7..3d571e16b 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate.go @@ -2,54 +2,183 @@ package ipfwdstate import ( "fmt" + "sync" log "github.com/sirupsen/logrus" "github.com/netbirdio/netbird/client/internal/routemanager/systemops" ) -// IPForwardingState is a struct that keeps track of the IP forwarding state. -// todo: read initial state of the IP forwarding from the system and reset the state based on it. -// todo: separate v4/v6 forwarding state, since the sysctls are independent -// (net.ipv4.ip_forward vs net.ipv6.conf.all.forwarding). Currently the nftables -// manager shares one instance between both routers, which works only because -// EnableIPForwarding enables both sysctls in a single call. +// IPForwardingState tracks v4 and v6 IP-forwarding sysctl enables with +// independent refcounts so a v4-only routing setup doesn't flip v6 sysctls. type IPForwardingState struct { - enabledCounter int + mu sync.Mutex + + v4Count int + v6Count int + + // routingV4/routingV6 track whether the routing path currently holds a + // reference, so repeated EnableRouting calls (one per network-map update) + // hold at most one reference per family and an unpaired DisableRouting + // can't release references held by DNAT rules. + routingV4 bool + routingV6 bool + + wgIfaceName string + v6Saved map[string]int } -func NewIPForwardingState() *IPForwardingState { - return &IPForwardingState{} +// NewIPForwardingState returns a state tracker for the IP-forwarding sysctls. +// wgIfaceName is excluded from the per-interface accept_ra handling. +func NewIPForwardingState(wgIfaceName string) *IPForwardingState { + return &IPForwardingState{wgIfaceName: wgIfaceName} } -func (f *IPForwardingState) RequestForwarding() error { - if f.enabledCounter != 0 { - f.enabledCounter++ +// Counts returns the current v4 and v6 refcounts. Intended for diagnostics +// and tests. +func (f *IPForwardingState) Counts() (v4, v6 int) { + f.mu.Lock() + defer f.mu.Unlock() + return f.v4Count, f.v6Count +} + +// RequestRouting takes the forwarding references for the routing path. It is +// idempotent: while routing already holds a reference, further calls don't +// increment the refcounts, and a v4-only request releases a previously held v6 +// reference. A v6 sysctl failure is logged and not returned so it can't take +// down v4 routing (the sysctl may be unwritable, e.g. read-only /proc/sys or +// IPv6 disabled on the kernel command line); v6 is retried on the next call. +func (f *IPForwardingState) RequestRouting(v6 bool) error { + f.mu.Lock() + defer f.mu.Unlock() + + if !f.routingV4 { + if err := f.requestV4(); err != nil { + return err + } + f.routingV4 = true + } + + if !v6 { + if !f.routingV6 { + return nil + } + f.routingV6 = false + return f.releaseV6() + } + + if f.routingV6 { return nil } - - if err := systemops.EnableIPForwarding(); err != nil { - return fmt.Errorf("failed to enable IP forwarding with sysctl: %w", err) + if err := f.requestV6(); err != nil { + log.Warnf("enable IPv6 forwarding for routing: %v", err) + return nil } - f.enabledCounter = 1 - log.Info("IP forwarding enabled") - + f.routingV6 = true return nil } -func (f *IPForwardingState) ReleaseForwarding() error { - if f.enabledCounter == 0 { - return nil +// ReleaseRouting releases the references RequestRouting holds. Calls without a +// held reference are no-ops. +func (f *IPForwardingState) ReleaseRouting() error { + f.mu.Lock() + defer f.mu.Unlock() + + if f.routingV4 { + f.routingV4 = false + f.releaseV4() } - - if f.enabledCounter > 1 { - f.enabledCounter-- - return nil + if f.routingV6 { + f.routingV6 = false + return f.releaseV6() } - - // if failed to disable IP forwarding we anyway decrement the counter - f.enabledCounter = 0 - - // todo call systemops.DisableIPForwarding() + return nil +} + +// RequestForwarding enables the family's forwarding sysctl on first request. +func (f *IPForwardingState) RequestForwarding(v6 bool) error { + f.mu.Lock() + defer f.mu.Unlock() + + if v6 { + return f.requestV6() + } + return f.requestV4() +} + +// ReleaseForwarding decrements the family counter. The last v6 release restores +// what enable captured. v4 stays on: net.ipv4.ip_forward is co-owned by other +// tooling (docker, k8s, libvirt). +func (f *IPForwardingState) ReleaseForwarding(v6 bool) error { + f.mu.Lock() + defer f.mu.Unlock() + + if v6 { + return f.releaseV6() + } + f.releaseV4() + return nil +} + +func (f *IPForwardingState) requestV4() error { + if f.v4Count == 0 { + if err := systemops.EnableV4IPForwarding(); err != nil { + return fmt.Errorf("enable IPv4 forwarding: %w", err) + } + log.Info("IPv4 forwarding enabled") + } + f.v4Count++ + return nil +} + +func (f *IPForwardingState) releaseV4() { + if f.v4Count > 0 { + f.v4Count-- + } +} + +func (f *IPForwardingState) requestV6() error { + if f.v6Count == 0 { + saved, err := systemops.EnableV6IPForwarding(f.wgIfaceName) + if err != nil { + if rerr := systemops.DisableV6IPForwarding(saved); rerr != nil { + log.Warnf("rollback partial v6 sysctls: %v", rerr) + } + return fmt.Errorf("enable IPv6 forwarding: %w", err) + } + // A failed restore on a previous release keeps its saved values; those + // are the true originals, so keep them over what this enable captured. + if f.v6Saved == nil { + f.v6Saved = saved + } else { + for k, v := range saved { + if _, ok := f.v6Saved[k]; !ok { + f.v6Saved[k] = v + } + } + } + log.Info("IPv6 forwarding enabled") + } + f.v6Count++ + return nil +} + +func (f *IPForwardingState) releaseV6() error { + if f.v6Count == 0 { + return nil + } + f.v6Count-- + if f.v6Count > 0 { + return nil + } + + // Keep the saved values on failure so a later release or enable/release + // cycle can still restore them; re-restoring an already-restored key is a + // no-op since the sysctl already holds the desired value. + if err := systemops.DisableV6IPForwarding(f.v6Saved); err != nil { + return fmt.Errorf("disable IPv6 forwarding: %w", err) + } + f.v6Saved = nil + log.Info("IPv6 forwarding disabled") return nil } diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go new file mode 100644 index 000000000..b4615ff02 --- /dev/null +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go @@ -0,0 +1,39 @@ +//go:build privileged + +package ipfwdstate + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request +// releases a previously held routing-owned v6 reference without touching +// references held by DNAT rules. +func TestRequestRoutingV6ToV4Transition(t *testing.T) { + f := NewIPForwardingState("wt-fwd-test") + + require.NoError(t, f.RequestRouting(true), "request routing with v6") + v4, v6 := f.Counts() + assert.Equal(t, 1, v4, "v4 reference held") + assert.Equal(t, 1, v6, "v6 reference held") + + require.NoError(t, f.RequestRouting(false), "request routing v4-only") + v4, v6 = f.Counts() + assert.Equal(t, 1, v4, "v4 reference kept") + assert.Equal(t, 0, v6, "routing-owned v6 reference released") + + // A DNAT-held reference survives a v4-only routing request. + require.NoError(t, f.RequestForwarding(true), "dnat v6 reference") + require.NoError(t, f.RequestRouting(false), "repeat v4-only request") + _, v6 = f.Counts() + assert.Equal(t, 1, v6, "dnat-held v6 reference survives") + require.NoError(t, f.ReleaseForwarding(true), "release dnat v6 reference") + + require.NoError(t, f.ReleaseRouting(), "release routing") + v4, v6 = f.Counts() + assert.Equal(t, 0, v4, "all v4 references released") + assert.Equal(t, 0, v6, "all v6 references released") +} diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go index 66b24cc5a..0ccfa83ac 100644 --- a/client/internal/routemanager/manager.go +++ b/client/internal/routemanager/manager.go @@ -8,14 +8,13 @@ import ( "net/netip" "net/url" "runtime" - "slices" + "sort" "strings" "sync" "sync/atomic" "syscall" "time" - "github.com/google/uuid" "github.com/hashicorp/go-multierror" log "github.com/sirupsen/logrus" "golang.org/x/exp/maps" @@ -52,15 +51,20 @@ type Manager interface { UpdateRoutes(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error ClassifyRoutes(newRoutes []*route.Route) (map[route.ID]*route.Route, route.HAMap) TriggerSelection(route.HAMap) + SelectRoutes(ids []route.NetID, appendRoute bool) error + DeselectRoutes(ids []route.NetID) error + SelectAllRoutes() + DeselectAllRoutes() GetRouteSelector() *routeselector.RouteSelector GetClientRoutes() route.HAMap GetSelectedClientRoutes() route.HAMap GetActiveClientRoutes() route.HAMap GetClientRoutesWithNetID() map[route.NetID][]*route.Route SetRouteChangeListener(listener listener.NetworkChangeListener) - InitialRouteRange() []string + CurrentRouteRange() []string SetFirewall(firewall.Manager) error SetDNSForwarderPort(port uint16) + ReconcilePeerAllowedIPs(peerKey string) error Stop(stateManager *statemanager.Manager) } @@ -71,10 +75,8 @@ type ManagerConfig struct { WGInterface iface.WGIface StatusRecorder *peer.Status RelayManager *relayClient.Manager - InitialRoutes []*route.Route StateManager *statemanager.Manager DNSServer dns.Server - DNSFeatureFlag bool PeerStore *peerstore.Store DisableClientRoutes bool DisableServerRoutes bool @@ -144,45 +146,12 @@ func NewManager(config ManagerConfig) *DefaultManager { useNoop := netstack.IsEnabled() || config.DisableClientRoutes dm.setupRefCounters(useNoop) - // don't proceed with client routes if it is disabled - if config.DisableClientRoutes { - return dm - } - - if runtime.GOOS == "android" { - dm.setupAndroidRoutes(config) - } return dm } -func (m *DefaultManager) setupAndroidRoutes(config ManagerConfig) { - cr := m.initialClientRoutes(config.InitialRoutes) - routesForComparison := slices.Clone(cr) - - if config.DNSFeatureFlag { - m.fakeIPManager = fakeip.NewManager() - - v4ID := uuid.NewString() - fakeIPRoute := &route.Route{ - ID: route.ID(v4ID), - Network: m.fakeIPManager.GetFakeIPBlock(), - NetID: route.NetID(v4ID), - Peer: m.pubKey, - NetworkType: route.IPv4Network, - } - v6ID := uuid.NewString() - fakeIPv6Route := &route.Route{ - ID: route.ID(v6ID), - Network: m.fakeIPManager.GetFakeIPv6Block(), - NetID: route.NetID(v6ID), - Peer: m.pubKey, - NetworkType: route.IPv6Network, - } - cr = append(cr, fakeIPRoute, fakeIPv6Route) - m.notifier.SetFakeIPRoutes([]*route.Route{fakeIPRoute, fakeIPv6Route}) - } - - m.notifier.SetInitialClientRoutes(cr, routesForComparison) +func (m *DefaultManager) enableFakeIPRoutes() { + m.fakeIPManager = fakeip.NewManager() + m.notifier.NotifyRouteChange() } func (m *DefaultManager) setupRefCounters(useNoop bool) { @@ -215,7 +184,7 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) { ) } - m.allowedIPsRefCounter = refcounter.New( + m.allowedIPsRefCounter = refcounter.NewAllowedIPs( func(prefix netip.Prefix, peerKey string) (string, error) { // save peerKey to use it in the remove function return peerKey, m.wgInterface.AddAllowedIP(peerKey, prefix) @@ -232,6 +201,30 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) { ) } +// ReconcilePeerAllowedIPs re-applies every routed allowed IP currently tracked for the peer +// onto the WireGuard device. The allowed-IP refcounter only calls its AddFunc (which pushes to +// the device) on a prefix's 0->1 transition, so a peer whose device entry was rebuilt without a +// matching refcounter change — e.g. a lazy connection cycling through idle->wake, which recreates +// the WireGuard peer with the overlay /32 only — ends up missing routed prefixes the refcounter +// still considers installed, and nothing retries. Calling this when the peer's WireGuard entry is +// (re)created restores convergence. It is add-only and idempotent: AddAllowedIP is update-only, so +// prefixes are re-added to an existing peer and an absent peer is left untouched. +func (m *DefaultManager) ReconcilePeerAllowedIPs(peerKey string) error { + if m.allowedIPsRefCounter == nil { + return nil + } + + return m.allowedIPsRefCounter.ReapplyMatching( + func(out string) bool { return out == peerKey }, + func(prefix netip.Prefix) error { + if err := m.wgInterface.AddAllowedIP(peerKey, prefix); err != nil { + return fmt.Errorf("add allowed IP %s for peer %s: %w", prefix, peerKey, err) + } + return nil + }, + ) +} + // Init sets up the routing func (m *DefaultManager) Init() error { m.routeSelector = m.initSelector() @@ -435,6 +428,9 @@ func (m *DefaultManager) UpdateRoutes( var merr *multierror.Error if !m.disableClientRoutes { + if runtime.GOOS == "android" && useNewDNSRoute && m.fakeIPManager == nil { + m.enableFakeIPRoutes() + } // Update route selector based on management server's isSelected status m.updateRouteSelectorFromManagement(clientRoutes) @@ -471,9 +467,32 @@ func (m *DefaultManager) SetRouteChangeListener(listener listener.NetworkChangeL m.notifier.SetListener(listener) } -// InitialRouteRange return the list of initial routes. It used by mobile systems -func (m *DefaultManager) InitialRouteRange() []string { - return m.notifier.GetInitialRouteRanges() +// CurrentRouteRange returns the current TUN route list. It is used by mobile systems +func (m *DefaultManager) CurrentRouteRange() []string { + m.mux.Lock() + defer m.mux.Unlock() + + if m.disableClientRoutes { + return nil + } + + filtered := m.routeSelector.FilterSelectedExitNodes(m.clientRoutes) + var nets []string + for _, routes := range filtered { + for _, r := range routes { + if r.IsDynamic() { + continue + } + nets = append(nets, r.NetString()) + } + } + + if m.fakeIPManager != nil { + nets = append(nets, m.fakeIPManager.GetFakeIPBlock().String(), m.fakeIPManager.GetFakeIPv6Block().String()) + } + + sort.Strings(nets) + return nets } // GetRouteSelector returns the route selector @@ -671,16 +690,6 @@ func (m *DefaultManager) ClassifyRoutes(newRoutes []*route.Route) (map[route.ID] return newServerRoutesMap, newClientRoutesIDMap } -func (m *DefaultManager) initialClientRoutes(initialRoutes []*route.Route) []*route.Route { - _, crMap := m.ClassifyRoutes(initialRoutes) - rs := make([]*route.Route, 0, len(crMap)) - for _, routes := range crMap { - rs = append(rs, routes...) - } - - return rs -} - func isRouteSupported(route *route.Route) bool { if netstack.IsEnabled() || !nbnet.CustomRoutingDisabled() || route.IsDynamic() { return true @@ -775,7 +784,7 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI var info exitNodeInfo for haID, routes := range clientRoutes { - if !m.isExitNodeRoute(routes) { + if !isExitNodeRoutes(routes) { continue } @@ -795,13 +804,6 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI return info } -func (m *DefaultManager) isExitNodeRoute(routes []*route.Route) bool { - if len(routes) == 0 { - return false - } - return route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network) -} - func (m *DefaultManager) categorizeUserSelection(netID route.NetID, info *exitNodeInfo) { if m.routeSelector.IsSelected(netID) { info.userSelected = append(info.userSelected, netID) diff --git a/client/internal/routemanager/mock.go b/client/internal/routemanager/mock.go index 937314995..2a8398b95 100644 --- a/client/internal/routemanager/mock.go +++ b/client/internal/routemanager/mock.go @@ -16,6 +16,8 @@ type MockManager struct { ClassifyRoutesFunc func(routes []*route.Route) (map[route.ID]*route.Route, route.HAMap) UpdateRoutesFunc func(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error TriggerSelectionFunc func(haMap route.HAMap) + SelectRoutesFunc func(ids []route.NetID, appendRoute bool) error + DeselectRoutesFunc func(ids []route.NetID) error GetRouteSelectorFunc func() *routeselector.RouteSelector GetClientRoutesFunc func() route.HAMap GetSelectedClientRoutesFunc func() route.HAMap @@ -28,8 +30,8 @@ func (m *MockManager) Init() error { return nil } -// InitialRouteRange mock implementation of InitialRouteRange from Manager interface -func (m *MockManager) InitialRouteRange() []string { +// CurrentRouteRange mock implementation of CurrentRouteRange from Manager interface +func (m *MockManager) CurrentRouteRange() []string { return nil } @@ -55,6 +57,30 @@ func (m *MockManager) TriggerSelection(networks route.HAMap) { } } +// SelectRoutes mock implementation of SelectRoutes from Manager interface +func (m *MockManager) SelectRoutes(ids []route.NetID, appendRoute bool) error { + if m.SelectRoutesFunc != nil { + return m.SelectRoutesFunc(ids, appendRoute) + } + return nil +} + +// DeselectRoutes mock implementation of DeselectRoutes from Manager interface +func (m *MockManager) DeselectRoutes(ids []route.NetID) error { + if m.DeselectRoutesFunc != nil { + return m.DeselectRoutesFunc(ids) + } + return nil +} + +// SelectAllRoutes mock implementation of SelectAllRoutes from Manager interface +func (m *MockManager) SelectAllRoutes() { +} + +// DeselectAllRoutes mock implementation of DeselectAllRoutes from Manager interface +func (m *MockManager) DeselectAllRoutes() { +} + // GetRouteSelector mock implementation of GetRouteSelector from Manager interface func (m *MockManager) GetRouteSelector() *routeselector.RouteSelector { if m.GetRouteSelectorFunc != nil { @@ -112,6 +138,11 @@ func (m *MockManager) SetFirewall(firewall.Manager) error { func (m *MockManager) SetDNSForwarderPort(port uint16) { } +// ReconcilePeerAllowedIPs mock implementation of ReconcilePeerAllowedIPs from Manager interface +func (m *MockManager) ReconcilePeerAllowedIPs(peerKey string) error { + return nil +} + // Stop mock implementation of Stop from Manager interface func (m *MockManager) Stop(stateManager *statemanager.Manager) { if m.StopFunc != nil { diff --git a/client/internal/routemanager/notifier/notifier_android.go b/client/internal/routemanager/notifier/notifier_android.go index 49300dbb2..5fa329310 100644 --- a/client/internal/routemanager/notifier/notifier_android.go +++ b/client/internal/routemanager/notifier/notifier_android.go @@ -6,7 +6,6 @@ import ( "net/netip" "slices" "sort" - "strings" "sync" "github.com/netbirdio/netbird/client/internal/listener" @@ -14,12 +13,15 @@ import ( ) type Notifier struct { - initialRoutes []*route.Route - currentRoutes []*route.Route - fakeIPRoutes []*route.Route + mu sync.Mutex - listener listener.NetworkChangeListener - listenerMux sync.Mutex + // currentRoutes is the last announced route set. It exists only to + // suppress noise: without it every network map sync would trigger the + // Java side, even when the routes did not change. The actual TUN route + // state is owned by the route manager and pulled from there. + currentRoutes []*route.Route + + listener listener.NetworkChangeListener } func NewNotifier() *Notifier { @@ -27,20 +29,15 @@ func NewNotifier() *Notifier { } func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { - n.listenerMux.Lock() - defer n.listenerMux.Unlock() + n.mu.Lock() + defer n.mu.Unlock() n.listener = listener } -// SetInitialClientRoutes stores the initial route sets for TUN configuration. -func (n *Notifier) SetInitialClientRoutes(initialRoutes []*route.Route, routesForComparison []*route.Route) { - n.initialRoutes = filterStatic(initialRoutes) - n.currentRoutes = filterStatic(routesForComparison) -} - -// SetFakeIPRoutes stores the fake IP routes to be included in every TUN rebuild. -func (n *Notifier) SetFakeIPRoutes(routes []*route.Route) { - n.fakeIPRoutes = routes +func (n *Notifier) NotifyRouteChange() { + n.mu.Lock() + defer n.mu.Unlock() + n.notifyLocked() } func (n *Notifier) OnNewRoutes(idMap route.HAMap) { @@ -54,46 +51,32 @@ func (n *Notifier) OnNewRoutes(idMap route.HAMap) { } } - if !n.hasRouteDiff(n.currentRoutes, newRoutes) { + n.mu.Lock() + defer n.mu.Unlock() + if !hasRouteDiff(n.currentRoutes, newRoutes) { return } n.currentRoutes = newRoutes - n.notify() + n.notifyLocked() } func (n *Notifier) OnNewPrefixes([]netip.Prefix) { // Not used on Android } -func (n *Notifier) notify() { - n.listenerMux.Lock() - defer n.listenerMux.Unlock() +func (n *Notifier) notifyLocked() { if n.listener == nil { return } - - allRoutes := slices.Clone(n.currentRoutes) - allRoutes = append(allRoutes, n.fakeIPRoutes...) - - routeStrings := n.routesToStrings(allRoutes) - sort.Strings(routeStrings) - go func(l listener.NetworkChangeListener) { - l.OnNetworkChanged(strings.Join(routeStrings, ",")) - }(n.listener) + n.listener.OnNetworkChanged("") } -func filterStatic(routes []*route.Route) []*route.Route { - out := make([]*route.Route, 0, len(routes)) - for _, r := range routes { - if !r.IsDynamic() { - out = append(out, r) - } - } - return out +func (n *Notifier) Close() { + // unused } -func (n *Notifier) routesToStrings(routes []*route.Route) []string { +func routesToStrings(routes []*route.Route) []string { nets := make([]string, 0, len(routes)) for _, r := range routes { nets = append(nets, r.NetString()) @@ -101,25 +84,10 @@ func (n *Notifier) routesToStrings(routes []*route.Route) []string { return nets } -func (n *Notifier) hasRouteDiff(a []*route.Route, b []*route.Route) bool { - slices.SortFunc(a, func(x, y *route.Route) int { - return strings.Compare(x.NetString(), y.NetString()) - }) - slices.SortFunc(b, func(x, y *route.Route) int { - return strings.Compare(x.NetString(), y.NetString()) - }) - - return !slices.EqualFunc(a, b, func(x, y *route.Route) bool { - return x.NetString() == y.NetString() - }) -} - -func (n *Notifier) GetInitialRouteRanges() []string { - initialStrings := n.routesToStrings(n.initialRoutes) - sort.Strings(initialStrings) - return initialStrings -} - -func (n *Notifier) Close() { - // unused +func hasRouteDiff(a []*route.Route, b []*route.Route) bool { + as := routesToStrings(a) + bs := routesToStrings(b) + sort.Strings(as) + sort.Strings(bs) + return !slices.Equal(as, bs) } diff --git a/client/internal/routemanager/notifier/notifier_ios.go b/client/internal/routemanager/notifier/notifier_ios.go index d0888f3a1..d663dd471 100644 --- a/client/internal/routemanager/notifier/notifier_ios.go +++ b/client/internal/routemanager/notifier/notifier_ios.go @@ -3,7 +3,6 @@ package notifier import ( - "container/list" "net/netip" "slices" "sort" @@ -16,20 +15,12 @@ import ( type Notifier struct { mu sync.Mutex - cond *sync.Cond currentPrefixes []string listener listener.NetworkChangeListener - queue *list.List - closed bool } func NewNotifier() *Notifier { - n := &Notifier{ - queue: list.New(), - } - n.cond = sync.NewCond(&n.mu) - go n.deliverLoop() - return n + return &Notifier{} } func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { @@ -38,11 +29,7 @@ func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { n.listener = listener } -func (n *Notifier) SetInitialClientRoutes([]*route.Route, []*route.Route) { - // iOS doesn't care about initial routes -} - -func (n *Notifier) SetFakeIPRoutes([]*route.Route) { +func (n *Notifier) NotifyRouteChange() { // Not used on iOS } @@ -59,44 +46,19 @@ func (n *Notifier) OnNewPrefixes(prefixes []netip.Prefix) { sort.Strings(newNets) n.mu.Lock() + defer n.mu.Unlock() if slices.Equal(n.currentPrefixes, newNets) { - n.mu.Unlock() return } n.currentPrefixes = newNets - routes := strings.Join(n.currentPrefixes, ",") - n.queue.PushBack(routes) - n.cond.Signal() - n.mu.Unlock() + if n.listener != nil { + n.listener.OnNetworkChanged(strings.Join(n.currentPrefixes, ",")) + } } func (n *Notifier) Close() { - n.mu.Lock() - n.closed = true - n.cond.Signal() - n.mu.Unlock() } func (n *Notifier) GetInitialRouteRanges() []string { return nil } - -func (n *Notifier) deliverLoop() { - for { - n.mu.Lock() - for n.queue.Len() == 0 && !n.closed { - n.cond.Wait() - } - if n.closed && n.queue.Len() == 0 { - n.mu.Unlock() - return - } - routes := n.queue.Remove(n.queue.Front()).(string) - l := n.listener - n.mu.Unlock() - - if l != nil { - l.OnNetworkChanged(routes) - } - } -} diff --git a/client/internal/routemanager/notifier/notifier_other.go b/client/internal/routemanager/notifier/notifier_other.go index 71b1096c2..fe48e07b3 100644 --- a/client/internal/routemanager/notifier/notifier_other.go +++ b/client/internal/routemanager/notifier/notifier_other.go @@ -19,11 +19,7 @@ func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { // Not used on non-mobile platforms } -func (n *Notifier) SetInitialClientRoutes([]*route.Route, []*route.Route) { - // Not used on non-mobile platforms -} - -func (n *Notifier) SetFakeIPRoutes([]*route.Route) { +func (n *Notifier) NotifyRouteChange() { // Not used on non-mobile platforms } @@ -35,10 +31,6 @@ func (n *Notifier) OnNewPrefixes(prefixes []netip.Prefix) { // Not used on non-mobile platforms } -func (n *Notifier) GetInitialRouteRanges() []string { - return []string{} -} - func (n *Notifier) Close() { // unused } diff --git a/client/internal/routemanager/reconcile_test.go b/client/internal/routemanager/reconcile_test.go new file mode 100644 index 000000000..c6806a6cd --- /dev/null +++ b/client/internal/routemanager/reconcile_test.go @@ -0,0 +1,90 @@ +//go:build !windows + +package routemanager + +import ( + "net" + "net/netip" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.zx2c4.com/wireguard/tun/netstack" + + "github.com/netbirdio/netbird/client/iface/device" + "github.com/netbirdio/netbird/client/iface/wgaddr" + "github.com/netbirdio/netbird/client/internal/routemanager/refcounter" +) + +// reconcileWGMock is a minimal iface.WGIface that only records AddAllowedIP calls; every other +// method is an inert stub because ReconcilePeerAllowedIPs exercises none of them. +type reconcileWGMock struct { + mu sync.Mutex + adds map[string][]netip.Prefix +} + +func (m *reconcileWGMock) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error { + m.mu.Lock() + defer m.mu.Unlock() + if m.adds == nil { + m.adds = map[string][]netip.Prefix{} + } + m.adds[peerKey] = append(m.adds[peerKey], allowedIP) + return nil +} + +func (m *reconcileWGMock) added(peerKey string) []netip.Prefix { + m.mu.Lock() + defer m.mu.Unlock() + return m.adds[peerKey] +} + +func (m *reconcileWGMock) RemoveAllowedIP(string, netip.Prefix) error { return nil } +func (m *reconcileWGMock) Name() string { return "utun-test" } +func (m *reconcileWGMock) Address() wgaddr.Address { return wgaddr.Address{} } +func (m *reconcileWGMock) ToInterface() *net.Interface { return nil } +func (m *reconcileWGMock) IsUserspaceBind() bool { return false } +func (m *reconcileWGMock) GetFilter() device.PacketFilter { return nil } +func (m *reconcileWGMock) GetDevice() *device.FilteredDevice { return nil } +func (m *reconcileWGMock) GetNet() *netstack.Net { return nil } + +// TestReconcilePeerAllowedIPs verifies the declarative reconcile re-applies every routed prefix +// tracked for the peer (self-heal, independent of refcount level) and stays scoped to that peer. +func TestReconcilePeerAllowedIPs(t *testing.T) { + wg := &reconcileWGMock{} + m := &DefaultManager{wgInterface: wg} + m.allowedIPsRefCounter = refcounter.NewAllowedIPs( + func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil }, + func(netip.Prefix, string) error { return nil }, + ) + + peerA1 := netip.MustParsePrefix("10.0.0.0/24") + peerA2 := netip.MustParsePrefix("10.1.0.0/24") + peerB1 := netip.MustParsePrefix("10.2.0.0/24") + + for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} { + _, err := m.allowedIPsRefCounter.Increment(prefix, peer) + require.NoError(t, err) + } + // Extra reference: reconcile must still re-apply the prefix even though its refcount never + // hit 0 again (the exact case the plain incremental path skips). + _, err := m.allowedIPsRefCounter.Increment(peerA1, "peerA") + require.NoError(t, err) + + require.NoError(t, m.ReconcilePeerAllowedIPs("peerA")) + + assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, wg.added("peerA"), + "reconcile must re-apply all routed prefixes of the peer") + assert.Empty(t, wg.added("peerB"), "reconcile must not touch another peer's prefixes") +} + +// TestReconcilePeerAllowedIPsNoCounter verifies reconcile is a safe no-op before the refcounter is +// set up. +func TestReconcilePeerAllowedIPsNoCounter(t *testing.T) { + wg := &reconcileWGMock{} + m := &DefaultManager{wgInterface: wg} + + require.NoError(t, m.ReconcilePeerAllowedIPs("peerA")) + assert.Empty(t, wg.added("peerA")) +} diff --git a/client/internal/routemanager/refcounter/allowedips.go b/client/internal/routemanager/refcounter/allowedips.go new file mode 100644 index 000000000..6d682e8a9 --- /dev/null +++ b/client/internal/routemanager/refcounter/allowedips.go @@ -0,0 +1,206 @@ +package refcounter + +import ( + "errors" + "fmt" + "net/netip" + "sort" + "sync" + + "github.com/hashicorp/go-multierror" + + nberrors "github.com/netbirdio/netbird/client/errors" +) + +// allowedIPsEntry holds the per-peer reference counts for a single prefix and which peer is +// currently installed in WireGuard. WireGuard allows a prefix on exactly one peer, so at most +// one peer is active at a time even when several peers reference the prefix. +type allowedIPsEntry struct { + // peers maps a peerKey to the number of references holding the prefix for that peer. + peers map[string]int + // active is the peerKey currently installed in WireGuard for this prefix ("" if none). + active string + // total is the sum of all per-peer reference counts (kept in sync with peers). + total int +} + +// AllowedIPsRefCounter is a peer-aware reference counter for WireGuard AllowedIPs. +// +// The generic Counter keys only by prefix and remembers a single Out value set by the first +// caller, which it never changes. That is wrong for AllowedIPs: two independent watchers (or +// multiple resolved domains) can reference the same prefix through different peers, and when the +// peer currently installed in WireGuard releases its last reference the prefix must be handed over +// to a surviving peer instead of being left pointing at the released one. +// +// It calls add/remove (which program WireGuard) only on the transitions that matter: +// - add on the first reference for a prefix, or when swapping the active peer; +// - remove on the last reference for a prefix, or on the old peer during a swap. +type AllowedIPsRefCounter struct { + mu sync.Mutex + entries map[netip.Prefix]*allowedIPsEntry + add AddFunc[netip.Prefix, string, string] + remove RemoveFunc[netip.Prefix, string] +} + +// NewAllowedIPs creates a new peer-aware AllowedIPs reference counter. +// add programs a prefix on a peer in WireGuard and returns the peerKey to store as the active peer. +// remove unprograms the prefix from the given peer. +func NewAllowedIPs(add AddFunc[netip.Prefix, string, string], remove RemoveFunc[netip.Prefix, string]) *AllowedIPsRefCounter { + return &AllowedIPsRefCounter{ + entries: map[netip.Prefix]*allowedIPsEntry{}, + add: add, + remove: remove, + } +} + +// Increment adds a reference to prefix for peerKey. WireGuard is programmed only for the first +// reference to a prefix; while a different peer is already installed the prefix is left with it +// (first peer wins, HA at the WireGuard layer is not possible) and only the reference count is kept. +func (rm *AllowedIPsRefCounter) Increment(prefix netip.Prefix, peerKey string) (Ref[string], error) { + rm.mu.Lock() + defer rm.mu.Unlock() + + e, ok := rm.entries[prefix] + if !ok { + e = &allowedIPsEntry{peers: map[string]int{}} + rm.entries[prefix] = e + } + + logCallerF("Increasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]", + prefix, peerKey, e.peers[peerKey], e.peers[peerKey]+1, e.total, e.total+1, e.active) + + // Program WireGuard only when nothing is installed yet for this prefix. + if e.active == "" { + out, err := rm.add(prefix, peerKey) + if errors.Is(err, ErrIgnore) { + if e.total == 0 { + delete(rm.entries, prefix) + } + return Ref[string]{Count: e.total, Out: e.active}, nil + } + if err != nil { + if e.total == 0 { + delete(rm.entries, prefix) + } + return Ref[string]{}, fmt.Errorf("failed to add allowed IP %v for peer %s: %w", prefix, peerKey, err) + } + e.active = out + } + + e.peers[peerKey]++ + e.total++ + + return Ref[string]{Count: e.total, Out: e.active}, nil +} + +// Decrement removes a reference to prefix for peerKey. When the peer currently installed in +// WireGuard releases its last reference, the prefix is swapped to a surviving peer if one exists, +// otherwise it is removed from WireGuard. +func (rm *AllowedIPsRefCounter) Decrement(prefix netip.Prefix, peerKey string) (Ref[string], error) { + rm.mu.Lock() + defer rm.mu.Unlock() + + e, ok := rm.entries[prefix] + if !ok { + logCallerF("No allowed IP reference found for prefix %v", prefix) + return Ref[string]{}, nil + } + + if e.peers[peerKey] > 0 { + logCallerF("Decreasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]", + prefix, peerKey, e.peers[peerKey], e.peers[peerKey]-1, e.total, e.total-1, e.active) + e.peers[peerKey]-- + e.total-- + if e.peers[peerKey] == 0 { + delete(e.peers, peerKey) + } + } else { + logCallerF("No allowed IP reference found for prefix %v peer %s", prefix, peerKey) + } + + // If the peer currently installed in WireGuard still holds references, nothing to reprogram. + // Keying the check on the active peer (not the one just released) makes this self-healing: + // a prior swap whose remove/add failed leaves e.active pointing at a peer with no references, + // and this retries the hand-off on the next Decrement instead of getting stuck. + if e.active != "" && e.peers[e.active] > 0 { + return Ref[string]{Count: e.total, Out: e.active}, nil + } + + // Detach the stale/gone active peer from WireGuard before reprogramming. + if e.active != "" { + if err := rm.remove(prefix, e.active); err != nil { + return Ref[string]{Count: e.total, Out: e.active}, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err) + } + e.active = "" + } + + // Hand the prefix over to a surviving peer, or drop the entry when none remain. + if survivor, ok := pickSurvivor(e.peers); ok { + out, err := rm.add(prefix, survivor) + if err != nil { + return Ref[string]{Count: e.total, Out: ""}, fmt.Errorf("swap allowed IP %v to peer %s: %w", prefix, survivor, err) + } + e.active = out + return Ref[string]{Count: e.total, Out: e.active}, nil + } + + delete(rm.entries, prefix) + return Ref[string]{Count: 0, Out: ""}, nil +} + +// Flush removes all prefixes from WireGuard and clears the counter. +func (rm *AllowedIPsRefCounter) Flush() error { + rm.mu.Lock() + defer rm.mu.Unlock() + + var merr *multierror.Error + for prefix, e := range rm.entries { + if e.active == "" { + continue + } + logCallerF("Flushing allowed IP for prefix %v peer %s", prefix, e.active) + if err := rm.remove(prefix, e.active); err != nil { + merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err)) + } + } + + clear(rm.entries) + + return nberrors.FormatErrorOrNil(merr) +} + +// ReapplyMatching calls apply for every prefix whose currently installed (active) peer satisfies +// pred, holding the lock for the whole pass. It is used to re-push allowed IPs onto a peer whose +// WireGuard entry was rebuilt (e.g. a lazy connection cycling idle->wake) without a matching +// refcounter change, which would otherwise leave the prefix installed in the counter but missing +// on the device. Only the active peer is considered — a prefix that lost its installed peer to a +// failed swap is skipped here and reconciled by the next Increment/Decrement. +func (rm *AllowedIPsRefCounter) ReapplyMatching(pred func(out string) bool, apply func(key netip.Prefix) error) error { + rm.mu.Lock() + defer rm.mu.Unlock() + + var merr *multierror.Error + for prefix, e := range rm.entries { + if e.active != "" && pred(e.active) { + if err := apply(prefix); err != nil { + merr = multierror.Append(merr, err) + } + } + } + return nberrors.FormatErrorOrNil(merr) +} + +// pickSurvivor deterministically selects a peer still referencing the prefix. WireGuard cannot do +// multipath for a single prefix, so any surviving peer is a valid winner; the choice is made stable +// (lowest peerKey) for predictable behavior and testability. +func pickSurvivor(peers map[string]int) (string, bool) { + if len(peers) == 0 { + return "", false + } + keys := make([]string, 0, len(peers)) + for k := range peers { + keys = append(keys, k) + } + sort.Strings(keys) + return keys[0], true +} diff --git a/client/internal/routemanager/refcounter/allowedips_test.go b/client/internal/routemanager/refcounter/allowedips_test.go new file mode 100644 index 000000000..835142083 --- /dev/null +++ b/client/internal/routemanager/refcounter/allowedips_test.go @@ -0,0 +1,241 @@ +package refcounter + +import ( + "errors" + "net/netip" + "testing" +) + +// fakeWG models WireGuard's cryptokey routing: a prefix can be installed on exactly one peer. +// failAdd/failRemove make the next add/remove fail once, to exercise the self-healing error paths. +type fakeWG struct { + installed map[netip.Prefix]string + adds int + removes int + failAdd bool + failRemove bool +} + +func newFakeWG() *fakeWG { + return &fakeWG{installed: map[netip.Prefix]string{}} +} + +func (f *fakeWG) counter() *AllowedIPsRefCounter { + return NewAllowedIPs( + func(prefix netip.Prefix, peerKey string) (string, error) { + if f.failAdd { + f.failAdd = false + return "", errors.New("add failed") + } + f.adds++ + f.installed[prefix] = peerKey + return peerKey, nil + }, + func(prefix netip.Prefix, peerKey string) error { + if f.failRemove { + f.failRemove = false + return errors.New("remove failed") + } + f.removes++ + // only clear if this peer is the one installed, mirroring wg semantics + if f.installed[prefix] == peerKey { + delete(f.installed, prefix) + } + return nil + }, + ) +} + +func mustPrefix(t *testing.T, s string) netip.Prefix { + t.Helper() + p, err := netip.ParsePrefix(s) + if err != nil { + t.Fatalf("parse prefix %q: %v", s, err) + } + return p +} + +func mustIncrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] { + t.Helper() + ref, err := c.Increment(p, peer) + if err != nil { + t.Fatalf("Increment(%v, %s): %v", p, peer, err) + } + return ref +} + +func mustDecrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] { + t.Helper() + ref, err := c.Decrement(p, peer) + if err != nil { + t.Fatalf("Decrement(%v, %s): %v", p, peer, err) + } + return ref +} + +// TestAllowedIPs_SwapOnActivePeerRemoval reproduces the reported bug: two networks with the same +// prefix routed by different peers. Removing the network whose peer is installed must hand the +// prefix over to the surviving peer instead of leaving it on the removed one. +func TestAllowedIPs_SwapOnActivePeerRemoval(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + // First peer wins while both are present. + if got := f.installed[p]; got != "peerA" { + t.Fatalf("expected peerA installed, got %q", got) + } + + // Remove the active peer's network -> must swap to peerB. + mustDecrement(t, c, p, "peerA") + if got := f.installed[p]; got != "peerB" { + t.Fatalf("BUG: prefix stuck on removed peer, want peerB got %q", got) + } + + // Remove the last one -> prefix gone. + mustDecrement(t, c, p, "peerB") + if _, ok := f.installed[p]; ok { + t.Fatalf("expected prefix removed, still installed on %q", f.installed[p]) + } +} + +// TestAllowedIPs_RemoveNonActivePeer removing a non-installed peer must not touch WireGuard. +func TestAllowedIPs_RemoveNonActivePeer(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + removesBefore := f.removes + + mustDecrement(t, c, p, "peerB") + if f.installed[p] != "peerA" { + t.Fatalf("active peer must stay peerA, got %q", f.installed[p]) + } + if f.removes != removesBefore { + t.Fatalf("removing a non-active peer must not call wg remove") + } +} + +// TestAllowedIPs_SamePeerMultipleRefs two references via the same peer must keep the prefix until +// the last reference is released (the reason the per-peer count must be an int, not a set). +func TestAllowedIPs_SamePeerMultipleRefs(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerA") + if f.adds != 1 { + t.Fatalf("expected a single wg add for the same peer, got %d", f.adds) + } + + mustDecrement(t, c, p, "peerA") + if f.installed[p] != "peerA" { + t.Fatalf("prefix must stay while a reference remains, got %q", f.installed[p]) + } + if f.removes != 0 { + t.Fatalf("no wg remove expected while a reference remains, got %d", f.removes) + } + + mustDecrement(t, c, p, "peerA") + if _, ok := f.installed[p]; ok { + t.Fatalf("prefix must be removed after last reference") + } +} + +// TestAllowedIPs_RefCountAndActive checks the Ref returned to callers (used for the HA-disabled log). +func TestAllowedIPs_RefCountAndActive(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + ref := mustIncrement(t, c, p, "peerA") + if ref.Count != 1 || ref.Out != "peerA" { + t.Fatalf("want {1, peerA}, got {%d, %q}", ref.Count, ref.Out) + } + ref = mustIncrement(t, c, p, "peerB") + if ref.Count != 2 || ref.Out != "peerA" { + t.Fatalf("want {2, peerA}, got {%d, %q}", ref.Count, ref.Out) + } +} + +// TestAllowedIPs_Flush removes everything installed and clears the counter. +func TestAllowedIPs_Flush(t *testing.T) { + f := newFakeWG() + c := f.counter() + p1 := mustPrefix(t, "10.44.8.0/24") + p2 := mustPrefix(t, "10.44.9.0/24") + + mustIncrement(t, c, p1, "peerA") + mustIncrement(t, c, p2, "peerB") + + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(f.installed) != 0 { + t.Fatalf("expected all prefixes removed, got %v", f.installed) + } + // After flush, a fresh increment must add again. + mustIncrement(t, c, p1, "peerC") + if f.installed[p1] != "peerC" { + t.Fatalf("counter not reset after flush") + } +} + +// TestAllowedIPs_SelfHealAfterSwapAddError ensures a failed add during a swap does not permanently +// strand the prefix: the next Decrement (or Increment) must retry and install a surviving peer. +func TestAllowedIPs_SelfHealAfterSwapAddError(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + mustIncrement(t, c, p, "peerC") + + // Removing the active peerA triggers a swap to a survivor; make the add fail once. + f.failAdd = true + if _, err := c.Decrement(p, "peerA"); err == nil { + t.Fatalf("expected error from failed swap add") + } + if _, ok := f.installed[p]; ok { + t.Fatalf("nothing should be installed after a failed swap add, got %q", f.installed[p]) + } + + // A later Decrement of a non-active survivor must retry the hand-off (self-heal), not stay stuck. + ref := mustDecrement(t, c, p, "peerC") + if got := f.installed[p]; got == "" { + t.Fatalf("self-heal failed: prefix left unrouted after add recovered") + } + if ref.Out == "" { + t.Fatalf("expected an active peer after self-heal, got empty") + } +} + +// TestAllowedIPs_SelfHealAfterRemoveError ensures a failed remove during a swap is retried instead +// of leaving e.active stuck on a peer that no longer holds references. +func TestAllowedIPs_SelfHealAfterRemoveError(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + + // Releasing active peerA must detach it (remove) then add peerB; fail the remove once. + f.failRemove = true + if _, err := c.Decrement(p, "peerA"); err == nil { + t.Fatalf("expected error from failed remove") + } + + // Next Decrement of the non-active survivor retries: removes stale peerA, installs peerB. + mustDecrement(t, c, p, "peerB") + // peerB had only one ref, so after retry the prefix is fully released. + if _, ok := f.installed[p]; ok { + t.Fatalf("expected prefix released after self-heal, still on %q", f.installed[p]) + } +} diff --git a/client/internal/routemanager/refcounter/refcounter.go b/client/internal/routemanager/refcounter/refcounter.go index 27a724f50..917120275 100644 --- a/client/internal/routemanager/refcounter/refcounter.go +++ b/client/internal/routemanager/refcounter/refcounter.go @@ -94,6 +94,26 @@ func (rm *Counter[Key, I, O]) Get(key Key) (Ref[O], bool) { return ref, ok } +// ReapplyMatching calls apply for every key whose stored Out satisfies pred, holding the +// counter lock for the whole pass. Running apply under the lock keeps it atomic with respect +// to Increment/Decrement: a prefix dropped to zero is removed from the map (and had its +// RemoveFunc called) before this pass observes it, so a stale key can never be re-applied. +// pred and apply are invoked under the lock, so they must not call back into the counter. +func (rm *Counter[Key, I, O]) ReapplyMatching(pred func(out O) bool, apply func(key Key) error) error { + rm.mu.Lock() + defer rm.mu.Unlock() + + var merr *multierror.Error + for key, ref := range rm.refCountMap { + if pred(ref.Out) { + if err := apply(key); err != nil { + merr = multierror.Append(merr, err) + } + } + } + return nberrors.FormatErrorOrNil(merr) +} + // Increment increments the reference count for the given key. // If this is the first reference to the key, the AddFunc is called. func (rm *Counter[Key, I, O]) Increment(key Key, in I) (Ref[O], error) { diff --git a/client/internal/routemanager/refcounter/refcounter_test.go b/client/internal/routemanager/refcounter/refcounter_test.go new file mode 100644 index 000000000..79a99c388 --- /dev/null +++ b/client/internal/routemanager/refcounter/refcounter_test.go @@ -0,0 +1,47 @@ +package refcounter + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// TestReapplyMatching verifies ReapplyMatching invokes apply for exactly the keys whose stored +// Out satisfies the predicate (no duplicates for multiply-referenced keys) — the primitive +// ReconcilePeerAllowedIPs relies on to re-apply a single peer's routed prefixes. +func TestReapplyMatching(t *testing.T) { + rc := New[netip.Prefix, string, string]( + func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil }, + func(netip.Prefix, string) error { return nil }, + ) + + peerA1 := netip.MustParsePrefix("10.0.0.0/24") + peerA2 := netip.MustParsePrefix("10.1.0.0/24") + peerB1 := netip.MustParsePrefix("10.2.0.0/24") + + for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} { + _, err := rc.Increment(prefix, peer) + require.NoError(t, err) + } + // a second reference must not make the key applied twice + _, err := rc.Increment(peerA1, "peerA") + require.NoError(t, err) + + var applied []netip.Prefix + err = rc.ReapplyMatching( + func(out string) bool { return out == "peerA" }, + func(key netip.Prefix) error { applied = append(applied, key); return nil }, + ) + require.NoError(t, err) + assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, applied) + + var none []netip.Prefix + err = rc.ReapplyMatching( + func(out string) bool { return out == "missing" }, + func(key netip.Prefix) error { none = append(none, key); return nil }, + ) + require.NoError(t, err) + assert.Empty(t, none) +} diff --git a/client/internal/routemanager/refcounter/types.go b/client/internal/routemanager/refcounter/types.go index aadac3e25..7da0e17e3 100644 --- a/client/internal/routemanager/refcounter/types.go +++ b/client/internal/routemanager/refcounter/types.go @@ -5,5 +5,7 @@ import "net/netip" // RouteRefCounter is a Counter for Route, it doesn't take any input on Increment and doesn't use any output on Decrement type RouteRefCounter = Counter[netip.Prefix, struct{}, struct{}] -// AllowedIPsRefCounter is a Counter for AllowedIPs, it takes a peer key on Increment and passes it back to Decrement -type AllowedIPsRefCounter = Counter[netip.Prefix, string, string] +// AllowedIPsRefCounter tracks WireGuard AllowedIPs per prefix. Unlike the generic Counter it is peer-aware: +// a prefix can be claimed by several peers at once and WireGuard allows a given prefix on exactly one peer, +// so the counter records the per-peer reference count and swaps the installed peer when the active one is released. +// See allowedips.go. diff --git a/client/internal/routemanager/selection.go b/client/internal/routemanager/selection.go new file mode 100644 index 000000000..6d5feec79 --- /dev/null +++ b/client/internal/routemanager/selection.go @@ -0,0 +1,138 @@ +package routemanager + +import ( + "fmt" + "slices" + + "github.com/hashicorp/go-multierror" + log "github.com/sirupsen/logrus" + "golang.org/x/exp/maps" + + nberrors "github.com/netbirdio/netbird/client/errors" + "github.com/netbirdio/netbird/route" +) + +// SelectRoutes selects the routes with the given network IDs and applies the +// new selection. V4/v6 exit-node pairs are expanded automatically. Exit nodes +// are mutually exclusive: if the selection activates an exit node, every other +// available exit node is deselected so two can't be active at once. With +// appendRoute=false the previous selection is replaced instead of extended. +func (m *DefaultManager) SelectRoutes(ids []route.NetID, appendRoute bool) error { + if err := m.selectRoutes(ids, appendRoute); err != nil { + return err + } + m.TriggerSelection(m.GetClientRoutes()) + return nil +} + +// DeselectRoutes removes the routes with the given network IDs from the +// selection and applies the change. V4/v6 exit-node pairs are expanded +// automatically. +func (m *DefaultManager) DeselectRoutes(ids []route.NetID) error { + if err := m.deselectRoutes(ids); err != nil { + return err + } + m.TriggerSelection(m.GetClientRoutes()) + return nil +} + +func (m *DefaultManager) deselectRoutes(ids []route.NetID) error { + routesMap := m.GetClientRoutesWithNetID() + routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap) + + log.Debugf("deselecting routes with ids: %v", routes) + + if err := m.routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil { + return fmt.Errorf("deselect routes: %w", err) + } + + return nil +} + +// SelectAllRoutes selects every available route and applies the selection. +// Exit nodes stay mutually exclusive: at most one remains active. +func (m *DefaultManager) SelectAllRoutes() { + m.selectAllRoutes() + m.TriggerSelection(m.GetClientRoutes()) +} + +func (m *DefaultManager) selectAllRoutes() { + m.routeSelector.SelectAllRoutes() + + // Select-all wipes every explicit selection, so exit nodes fall back to + // management's auto-apply flags — which may mark several at once. + // Reconcile immediately so at most one exit node stays active instead of + // waiting for the next network map to enforce it. + m.mux.Lock() + defer m.mux.Unlock() + m.updateRouteSelectorFromManagement(m.clientRoutes) +} + +// DeselectAllRoutes deselects every route and applies the change. +func (m *DefaultManager) DeselectAllRoutes() { + m.routeSelector.DeselectAllRoutes() + m.TriggerSelection(m.GetClientRoutes()) +} + +func (m *DefaultManager) selectRoutes(ids []route.NetID, appendRoute bool) error { + routesMap := m.GetClientRoutesWithNetID() + routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap) + allIDs := maps.Keys(routesMap) + + log.Debugf("selecting routes with ids: %v", routes) + + // A partial failure (e.g. an unknown ID in the request) still selects the + // valid routes, so exclusivity below must run regardless of the error. + var merr *multierror.Error + if err := m.routeSelector.SelectRoutes(routes, appendRoute, allIDs); err != nil { + merr = multierror.Append(merr, fmt.Errorf("select routes: %w", err)) + } + + // Exit nodes are mutually exclusive: if this selection activates an + // exit node, deselect every other available exit node so two can't be + // selected at once. Non-exit route selections are left untouched. + if requestActivatesExitNode(routes, routesMap) { + if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 { + if err := m.routeSelector.DeselectRoutes(others, allIDs); err != nil { + merr = multierror.Append(merr, fmt.Errorf("deselect sibling exit nodes: %w", err)) + } + } + } + + return nberrors.FormatErrorOrNil(merr) +} + +func isExitNodeRoutes(routes []*route.Route) bool { + return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network)) +} + +// requestActivatesExitNode reports whether any requested NetID maps to an exit +// node (default route) in the current route table. +func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool { + for _, id := range requested { + if isExitNodeRoutes(routesMap[id]) { + return true + } + } + return false +} + +// otherExitNodeIDs returns every available exit-node NetID that is not in the +// requested set — the siblings to deselect so a single exit node stays active. +func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID { + keep := make(map[route.NetID]struct{}, len(requested)) + for _, id := range requested { + keep[id] = struct{}{} + } + var others []route.NetID + for id, routes := range routesMap { + if !isExitNodeRoutes(routes) { + continue + } + if _, ok := keep[id]; ok { + continue + } + others = append(others, id) + } + return others +} diff --git a/client/internal/routemanager/selection_test.go b/client/internal/routemanager/selection_test.go new file mode 100644 index 000000000..6066b5661 --- /dev/null +++ b/client/internal/routemanager/selection_test.go @@ -0,0 +1,129 @@ +package routemanager + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/routeselector" + "github.com/netbirdio/netbird/route" +) + +func v6ExitRoute(netID, peer string) *route.Route { + return &route.Route{ + NetID: route.NetID(netID), + Network: netip.MustParsePrefix("::/0"), + Peer: peer, + } +} + +func newSelectionTestManager() *DefaultManager { + return &DefaultManager{ + routeSelector: routeselector.NewRouteSelector(), + clientRoutes: route.HAMap{ + "exitA|0.0.0.0/0": {exitRoute("exitA", "p1", true)}, + "exitA-v6|::/0": {v6ExitRoute("exitA-v6", "p1")}, + "exitB|0.0.0.0/0": {exitRoute("exitB", "p2", true)}, + "lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}}, + }, + } +} + +func TestSelectRoutes_ExitNodeExclusivity(t *testing.T) { + m := newSelectionTestManager() + + // Selecting an exit node selects its v6 pair and deselects the sibling. + require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true)) + assert.True(t, m.routeSelector.IsSelected("exitA"), "exitA should be selected") + assert.True(t, m.routeSelector.IsSelected("exitA-v6"), "the v6 pair follows its v4 base") + assert.False(t, m.routeSelector.IsSelected("exitB"), "the sibling exit node must be deselected") + + // Switching to the sibling deselects the previous exit node and its v6 pair. + require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true)) + assert.True(t, m.routeSelector.IsSelected("exitB"), "exitB should now be selected") + assert.False(t, m.routeSelector.IsSelected("exitA"), "the previous exit node must be deselected") + assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "the previous exit node's v6 pair must be deselected") + assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched") + + // Selecting a non-exit route leaves the active exit node alone. + require.NoError(t, m.selectRoutes([]route.NetID{"lan"}, true)) + assert.True(t, m.routeSelector.IsSelected("exitB"), "selecting a non-exit route keeps the exit node") + + // Deselecting the active exit node turns every exit node off. + require.NoError(t, m.deselectRoutes([]route.NetID{"exitB"})) + assert.False(t, m.routeSelector.IsSelected("exitB"), "exitB should be deselected") + assert.False(t, m.routeSelector.IsSelected("exitA"), "exitA stays deselected") + assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched") +} + +func TestSelectRoutes_PartialErrorStillEnforcesExclusivity(t *testing.T) { + // The unknown ID must be reported, but the valid exit node in the same + // request is still selected — so its sibling must still be deselected. + // Both orderings are covered: processing must continue past the invalid + // ID wherever it sits in the request. + requests := map[string][]route.NetID{ + "invalid id first": {"missing", "exitB"}, + "invalid id last": {"exitB", "missing"}, + } + + for name, ids := range requests { + t.Run(name, func(t *testing.T) { + m := newSelectionTestManager() + + require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true)) + + err := m.selectRoutes(ids, true) + assert.Error(t, err, "unknown id must be reported") + assert.True(t, m.routeSelector.IsSelected("exitB"), "valid exit node from the request is selected") + assert.False(t, m.routeSelector.IsSelected("exitA"), "sibling exit node must be deselected despite the error") + assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "sibling's v6 pair must be deselected too") + }) + } +} + +func TestSelectAllRoutes_KeepsSingleExitNode(t *testing.T) { + // Both exit nodes are marked for auto-apply by management + // (SkipAutoApply=false), the state where select-all could turn on two at + // once without the immediate reconciliation. + m := &DefaultManager{ + routeSelector: routeselector.NewRouteSelector(), + clientRoutes: route.HAMap{ + "exitA|0.0.0.0/0": {exitRoute("exitA", "p1", false)}, + "exitB|0.0.0.0/0": {exitRoute("exitB", "p2", false)}, + "lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}}, + }, + } + + require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true)) + + m.selectAllRoutes() + + assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit routes are all selected") + assert.True(t, m.routeSelector.IsSelected("exitA"), "the deterministic management pick stays active") + assert.False(t, m.routeSelector.IsSelected("exitB"), "select-all must not leave a second exit node active") +} + +func TestSelectRoutes_UnknownRoute(t *testing.T) { + m := newSelectionTestManager() + + assert.Error(t, m.selectRoutes([]route.NetID{"missing"}, true), "selecting an unavailable route must fail") + assert.Error(t, m.deselectRoutes([]route.NetID{"missing"}), "deselecting an unavailable route must fail") +} + +func TestExitNodeSelectionHelpers(t *testing.T) { + routesMap := map[route.NetID][]*route.Route{ + "exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}}, + "exitB": {{Network: netip.MustParsePrefix("::/0")}}, + "lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}}, + } + + assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node") + assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node") + assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node") + assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node") + + others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"}) + assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored") +} diff --git a/client/internal/routemanager/static/route.go b/client/internal/routemanager/static/route.go index d480fdf00..8ba03d090 100644 --- a/client/internal/routemanager/static/route.go +++ b/client/internal/routemanager/static/route.go @@ -15,6 +15,11 @@ type Route struct { route *route.Route routeRefCounter *refcounter.RouteRefCounter allowedIPsRefcounter *refcounter.AllowedIPsRefCounter + // currentPeerKey is the routing peer this watcher currently has the prefix installed on + // (the HA winner elected by the watcher). It can differ from route.Peer and change on + // failover, so it is recorded on AddAllowedIPs and used on RemoveAllowedIPs to decrement + // the exact peer that was incremented. + currentPeerKey string } func NewRoute(params common.HandlerParams) *Route { @@ -52,12 +57,15 @@ func (r *Route) AddAllowedIPs(peerKey string) error { ref.Out, ) } + r.currentPeerKey = peerKey return nil } func (r *Route) RemoveAllowedIPs() error { - if _, err := r.allowedIPsRefcounter.Decrement(r.route.Network); err != nil { - return err + var err error + if _, decErr := r.allowedIPsRefcounter.Decrement(r.route.Network, r.currentPeerKey); decErr != nil { + err = fmt.Errorf("remove allowed IP %s: %w", r.route.Network, decErr) } - return nil + r.currentPeerKey = "" + return err } diff --git a/client/internal/routemanager/sysctl/sysctl_linux.go b/client/internal/routemanager/sysctl/sysctl_linux.go index f96a57f37..bb131c691 100644 --- a/client/internal/routemanager/sysctl/sysctl_linux.go +++ b/client/internal/routemanager/sysctl/sysctl_linux.go @@ -20,6 +20,8 @@ const ( rpFilterPath = "net.ipv4.conf.all.rp_filter" rpFilterInterfacePath = "net.ipv4.conf.%s.rp_filter" srcValidMarkPath = "net.ipv4.conf.all.src_valid_mark" + percentEscape = "%25" + dotEscape = "%2E" ) type iface interface { @@ -56,7 +58,7 @@ func Setup(wgIface iface) (map[string]int, error) { continue } - i := fmt.Sprintf(rpFilterInterfacePath, intf.Name) + i := fmt.Sprintf(rpFilterInterfacePath, EscapeInterfaceName(intf.Name)) oldVal, err := Set(i, 2, true) if err != nil { result = multierror.Append(result, err) @@ -68,9 +70,20 @@ func Setup(wgIface iface) (map[string]int, error) { return keys, nberrors.FormatErrorOrNil(result) } +// EscapeInterfaceName escapes '%' and '.' in an interface name (e.g. VLANs +// like eth0.100) so the name survives the dot-to-slash conversion in Set. +func EscapeInterfaceName(name string) string { + safe := strings.ReplaceAll(name, "%", percentEscape) + return strings.ReplaceAll(safe, ".", dotEscape) +} + // Set sets a sysctl configuration, if onlyIfOne is true it will only set the new value if it's set to 1 func Set(key string, desiredValue int, onlyIfOne bool) (int, error) { - path := fmt.Sprintf("/proc/sys/%s", strings.ReplaceAll(key, ".", "/")) + path := strings.ReplaceAll(key, ".", "/") + // Unescape interface dots and percent signs + path = strings.ReplaceAll(path, dotEscape, ".") + path = strings.ReplaceAll(path, percentEscape, "%") + path = fmt.Sprintf("/proc/sys/%s", path) currentValue, err := os.ReadFile(path) if err != nil { return -1, fmt.Errorf("read sysctl %s: %w", key, err) diff --git a/client/internal/routemanager/systemops/routeselection_windows_test.go b/client/internal/routemanager/systemops/routeselection_windows_test.go new file mode 100644 index 000000000..108338dd9 --- /dev/null +++ b/client/internal/routemanager/systemops/routeselection_windows_test.go @@ -0,0 +1,82 @@ +//go:build windows + +package systemops + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestSortRouteCandidates(t *testing.T) { + tests := []struct { + name string + candidates []candidateRoute + wantOrder []uint32 + }{ + { + name: "longest prefix wins over metrics", + candidates: []candidateRoute{ + {interfaceIndex: 1, prefixLength: 0, routeMetric: 0, interfaceMetric: 5}, + {interfaceIndex: 2, prefixLength: 24, routeMetric: 100, interfaceMetric: 50}, + }, + wantOrder: []uint32{2, 1}, + }, + { + // Windows ranks equal-length prefixes by route metric + interface metric, + // so a higher route metric on a low metric interface can still win. + name: "combined metric beats route metric alone", + candidates: []candidateRoute{ + {interfaceIndex: 8, prefixLength: 0, routeMetric: 0, interfaceMetric: 100}, + {interfaceIndex: 5, prefixLength: 0, routeMetric: 10, interfaceMetric: 5}, + }, + wantOrder: []uint32{5, 8}, + }, + { + name: "lower combined metric wins", + candidates: []candidateRoute{ + {interfaceIndex: 5, prefixLength: 0, routeMetric: 300, interfaceMetric: 5}, + {interfaceIndex: 8, prefixLength: 0, routeMetric: 0, interfaceMetric: 100}, + }, + wantOrder: []uint32{8, 5}, + }, + { + name: "equal combined metric falls back to route metric", + candidates: []candidateRoute{ + {interfaceIndex: 1, prefixLength: 0, routeMetric: 20, interfaceMetric: 10}, + {interfaceIndex: 2, prefixLength: 0, routeMetric: 5, interfaceMetric: 25}, + }, + wantOrder: []uint32{2, 1}, + }, + { + // The metrics are uint32 on the Windows side, so the sum must not wrap. + name: "combined metric beyond the uint32 range", + candidates: []candidateRoute{ + {interfaceIndex: 1, prefixLength: 0, routeMetric: math.MaxUint32, interfaceMetric: 5}, + {interfaceIndex: 2, prefixLength: 0, routeMetric: math.MaxUint32 - 10, interfaceMetric: 5}, + }, + wantOrder: []uint32{2, 1}, + }, + { + name: "unknown interface metric ranks on route metric only", + candidates: []candidateRoute{ + {interfaceIndex: 1, prefixLength: 0, routeMetric: 30, interfaceMetric: -1}, + {interfaceIndex: 2, prefixLength: 0, routeMetric: 5, interfaceMetric: 10}, + }, + wantOrder: []uint32{2, 1}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + sortRouteCandidates(tt.candidates) + + got := make([]uint32, 0, len(tt.candidates)) + for _, c := range tt.candidates { + got = append(got, c.interfaceIndex) + } + assert.Equal(t, tt.wantOrder, got) + }) + } +} diff --git a/client/internal/routemanager/systemops/systemops_android.go b/client/internal/routemanager/systemops/systemops_android.go index 7cb8dae93..97b4ed8ec 100644 --- a/client/internal/routemanager/systemops/systemops_android.go +++ b/client/internal/routemanager/systemops/systemops_android.go @@ -32,8 +32,17 @@ func (r *SysOps) removeFromRouteTable(netip.Prefix, Nexthop) error { return nil } -func EnableIPForwarding() error { - log.Infof("Enable IP forwarding is not implemented on %s", runtime.GOOS) +func EnableV4IPForwarding() error { + log.Infof("Enable IPv4 forwarding is not implemented on %s", runtime.GOOS) + return nil +} + +func EnableV6IPForwarding(string) (map[string]int, error) { + log.Infof("Enable IPv6 forwarding is not implemented on %s", runtime.GOOS) + return map[string]int{}, nil +} + +func DisableV6IPForwarding(map[string]int) error { return nil } diff --git a/client/internal/routemanager/systemops/systemops_ios.go b/client/internal/routemanager/systemops/systemops_ios.go index 99a363371..0cccd4962 100644 --- a/client/internal/routemanager/systemops/systemops_ios.go +++ b/client/internal/routemanager/systemops/systemops_ios.go @@ -58,8 +58,17 @@ func (r *SysOps) removeFromRouteTable(netip.Prefix, Nexthop) error { return nil } -func EnableIPForwarding() error { - log.Infof("Enable IP forwarding is not implemented on %s", runtime.GOOS) +func EnableV4IPForwarding() error { + log.Infof("Enable IPv4 forwarding is not implemented on %s", runtime.GOOS) + return nil +} + +func EnableV6IPForwarding(string) (map[string]int, error) { + log.Infof("Enable IPv6 forwarding is not implemented on %s", runtime.GOOS) + return map[string]int{}, nil +} + +func DisableV6IPForwarding(map[string]int) error { return nil } diff --git a/client/internal/routemanager/systemops/systemops_linux.go b/client/internal/routemanager/systemops/systemops_linux.go index 8c6b7d9a9..7d608d886 100644 --- a/client/internal/routemanager/systemops/systemops_linux.go +++ b/client/internal/routemanager/systemops/systemops_linux.go @@ -763,13 +763,10 @@ func flushRoutes(tableID, family int) error { return nberrors.FormatErrorOrNil(result) } -func EnableIPForwarding() error { +func EnableV4IPForwarding() error { if _, err := sysctl.Set(ipv4ForwardingPath, 1, false); err != nil { return err } - if _, err := sysctl.Set(ipv6ForwardingPath, 1, false); err != nil { - log.Warnf("failed to enable IPv6 forwarding: %v", err) - } return nil } diff --git a/client/internal/routemanager/systemops/systemops_nonlinux.go b/client/internal/routemanager/systemops/systemops_nonlinux.go index 016a62ebd..837ac0cd2 100644 --- a/client/internal/routemanager/systemops/systemops_nonlinux.go +++ b/client/internal/routemanager/systemops/systemops_nonlinux.go @@ -43,8 +43,17 @@ func (r *SysOps) RemoveVPNRoute(prefix netip.Prefix, intf *net.Interface) error return r.genericRemoveVPNRoute(prefix, intf) } -func EnableIPForwarding() error { - log.Infof("Enable IP forwarding is not implemented on %s", runtime.GOOS) +func EnableV4IPForwarding() error { + log.Infof("Enable IPv4 forwarding is not implemented on %s", runtime.GOOS) + return nil +} + +func EnableV6IPForwarding(string) (map[string]int, error) { + log.Infof("Enable IPv6 forwarding is not implemented on %s", runtime.GOOS) + return map[string]int{}, nil +} + +func DisableV6IPForwarding(map[string]int) error { return nil } diff --git a/client/internal/routemanager/systemops/systemops_windows.go b/client/internal/routemanager/systemops/systemops_windows.go index 7bce6af80..47d556cf6 100644 --- a/client/internal/routemanager/systemops/systemops_windows.go +++ b/client/internal/routemanager/systemops/systemops_windows.go @@ -882,26 +882,40 @@ func getInterfaceMetric(interfaceIndex uint32, family int16) int { return int(ipInterfaceRow.Metric) } -// sortRouteCandidates sorts route candidates by priority: prefix length -> route metric -> interface metric +// sortRouteCandidates sorts route candidates by priority: prefix length -> combined metric -> route metric. +// Windows prefers the longest matching prefix and, among prefixes of the same length, the lowest metric, see +// https://learn.microsoft.com/en-us/windows-hardware/customize/desktop/unattend/microsoft-windows-tcpip-interfaces-interface-routes-route-metric func sortRouteCandidates(candidates []candidateRoute) { sort.Slice(candidates, func(i, j int) bool { if candidates[i].prefixLength != candidates[j].prefixLength { return candidates[i].prefixLength > candidates[j].prefixLength } - if candidates[i].routeMetric != candidates[j].routeMetric { - return candidates[i].routeMetric < candidates[j].routeMetric + mi, mj := combinedMetric(candidates[i]), combinedMetric(candidates[j]) + if mi != mj { + return mi < mj } - return candidates[i].interfaceMetric < candidates[j].interfaceMetric + return candidates[i].routeMetric < candidates[j].routeMetric }) } +// combinedMetric returns the effective metric Windows uses to rank routes with an equal prefix length: +// the sum of the route metric and the metric of the interface the route is on, see +// https://learn.microsoft.com/en-us/windows-server/networking/technologies/network-subsystem/net-sub-interface-metric +// An unknown interface metric contributes nothing. +func combinedMetric(candidate candidateRoute) uint64 { + if candidate.interfaceMetric < 0 { + return uint64(candidate.routeMetric) + } + return uint64(candidate.routeMetric) + uint64(candidate.interfaceMetric) +} + // GetBestInterface finds the best interface for reaching a destination, // excluding the VPN interface to avoid routing loops. // // Route selection priority: // 1. Longest prefix match (most specific route) -// 2. Lowest route metric -// 3. Lowest interface metric +// 2. Lowest combined metric (route metric + interface metric) +// 3. Lowest route metric. func GetBestInterface(dest netip.Addr, vpnIntf string) (*net.Interface, error) { var skipInterfaceIndex int if vpnIntf != "" { @@ -925,7 +939,6 @@ func GetBestInterface(dest netip.Addr, vpnIntf string) (*net.Interface, error) { return nil, fmt.Errorf("no route to %s", dest) } - // Sort routes: prefix length -> route metric -> interface metric sortRouteCandidates(candidates) for _, candidate := range candidates { diff --git a/client/internal/routemanager/systemops/v6forwarding_linux.go b/client/internal/routemanager/systemops/v6forwarding_linux.go new file mode 100644 index 000000000..c1e0d4588 --- /dev/null +++ b/client/internal/routemanager/systemops/v6forwarding_linux.go @@ -0,0 +1,92 @@ +//go:build !android + +package systemops + +import ( + "fmt" + "net" + "os" + + "github.com/hashicorp/go-multierror" + log "github.com/sirupsen/logrus" + + nberrors "github.com/netbirdio/netbird/client/errors" + "github.com/netbirdio/netbird/client/internal/routemanager/sysctl" +) + +const ( + // 1 (default) accepts RAs only while forwarding is off; 2 keeps RA + // acceptance on regardless, so RA-installed host defaults survive our + // v6 forwarding flip. + acceptRAInterfacePath = "net.ipv6.conf.%s.accept_ra" + acceptRADefaultPath = "net.ipv6.conf.default.accept_ra" + acceptRAProcPathFormat = "/proc/sys/net/ipv6/conf/%s/accept_ra" +) + +// EnableV6IPForwarding bumps accept_ra=2 on host v6 interfaces before flipping +// forwarding=1, so RA-installed host defaults survive. Returns the prior values +// of sysctls we actually changed; entries already at the target are omitted. +func EnableV6IPForwarding(wgIfaceName string) (map[string]int, error) { + saved := map[string]int{} + bumpAcceptRA(saved, wgIfaceName) + + oldVal, err := sysctl.Set(ipv6ForwardingPath, 1, false) + if err != nil { + return saved, err + } + if oldVal != 1 { + saved[ipv6ForwardingPath] = oldVal + } + return saved, nil +} + +// DisableV6IPForwarding restores what EnableV6IPForwarding captured. +func DisableV6IPForwarding(saved map[string]int) error { + var result *multierror.Error + for key, value := range saved { + if _, err := sysctl.Set(key, value, false); err != nil { + result = multierror.Append(result, fmt.Errorf("restore %s: %w", key, err)) + } + } + return nberrors.FormatErrorOrNil(result) +} + +func bumpAcceptRA(saved map[string]int, wgIfaceName string) { + // Also bump conf.default so interfaces created while forwarding is on + // (hotplug, new Wi-Fi/dock) inherit accept_ra=2 and keep accepting RAs. + bumpAcceptRAKey(saved, acceptRADefaultPath) + + interfaces, err := net.Interfaces() + if err != nil { + log.Warnf("list interfaces for accept_ra: %v", err) + return + } + for _, intf := range interfaces { + if intf.Name == "lo" || intf.Name == wgIfaceName { + continue + } + bumpAcceptRAForInterface(saved, intf.Name) + } +} + +func bumpAcceptRAForInterface(saved map[string]int, name string) { + // Build procfs path from name, not the dotted key: VLAN names like eth0.100. + if _, err := os.Stat(fmt.Sprintf(acceptRAProcPathFormat, name)); err != nil { + return + } + bumpAcceptRAKey(saved, fmt.Sprintf(acceptRAInterfacePath, sysctl.EscapeInterfaceName(name))) +} + +func bumpAcceptRAKey(saved map[string]int, key string) { + // onlyIfOne=true: leave admin overrides (0, 2) alone. + oldVal, err := sysctl.Set(key, 2, true) + if err != nil { + log.Warnf("bump %s: %v", key, err) + return + } + // With onlyIfOne, a write only happened when the old value was 1; values + // left untouched (0, 2) must not be recorded for restore. + if oldVal == 1 { + saved[key] = oldVal + } +} diff --git a/client/internal/routemanager/systemops/v6route_linux_test.go b/client/internal/routemanager/systemops/v6route_linux_test.go index 449d4cbd2..d8c0012d1 100644 --- a/client/internal/routemanager/systemops/v6route_linux_test.go +++ b/client/internal/routemanager/systemops/v6route_linux_test.go @@ -5,6 +5,7 @@ package systemops import ( "errors" "net" + "net/netip" "syscall" "testing" @@ -29,6 +30,7 @@ func ensureIPv6DefaultRoute(t *testing.T) { } if err := netlink.RouteAdd(route); err != nil { if errors.Is(err, syscall.EEXIST) { + requireUsableIPv6Nexthop(t) return } t.Skipf("install IPv6 fallback default route: %v", err) @@ -38,4 +40,36 @@ func ensureIPv6DefaultRoute(t *testing.T) { t.Logf("delete IPv6 fallback default route: %v", err) } }) + + requireUsableIPv6Nexthop(t) +} + +// requireUsableIPv6Nexthop skips the test unless the resolved IPv6 default +// nexthop can actually carry a route. Installing the default route succeeding +// does not imply the kernel accepts it as a nexthop for a concrete prefix. +func requireUsableIPv6Nexthop(t *testing.T) { + t.Helper() + + nexthop, err := GetNextHop(netip.IPv6Unspecified()) + if err != nil { + t.Skipf("resolve IPv6 default nexthop: %v", err) + } + + probe := &netlink.Route{ + Scope: netlink.SCOPE_UNIVERSE, + Table: syscall.RT_TABLE_MAIN, + Family: netlink.FAMILY_V6, + Dst: &net.IPNet{IP: net.ParseIP("100::64"), Mask: net.CIDRMask(128, 128)}, + } + require.NoError(t, addNextHop(nexthop, probe), "build IPv6 probe route") + + switch err := netlink.RouteAdd(probe); { + case err == nil: + if err := netlink.RouteDel(probe); err != nil && !errors.Is(err, syscall.ESRCH) { + t.Logf("delete IPv6 probe route: %v", err) + } + case errors.Is(err, syscall.EEXIST): + default: + t.Skipf("IPv6 nexthop %s unusable for route installation: %v", nexthop, err) + } } diff --git a/client/internal/sleep/service.go b/client/internal/sleep/service.go index 196a33f52..93691c4c7 100644 --- a/client/internal/sleep/service.go +++ b/client/internal/sleep/service.go @@ -18,8 +18,8 @@ type Service struct { } func New() (*Service, error) { - d, err := NewDetector() - if err != nil { + d, err := NewDetector() //nolint:staticcheck + if err != nil { //nolint:staticcheck // always errors on platforms without a sleep detector return nil, err } diff --git a/client/internal/tunnelnotifier/notifier.go b/client/internal/tunnelnotifier/notifier.go new file mode 100644 index 000000000..b62923a6e --- /dev/null +++ b/client/internal/tunnelnotifier/notifier.go @@ -0,0 +1,124 @@ +package tunnelnotifier + +import ( + "container/list" + "sync" + + "github.com/netbirdio/netbird/client/internal/dns" + "github.com/netbirdio/netbird/client/internal/listener" +) + +type eventKind int + +const ( + eventRoutes eventKind = iota + eventIfaceIP + eventIfaceIPv6 + eventDNS +) + +var ( + _ listener.NetworkChangeListener = (*Notifier)(nil) + _ dns.IosDnsManager = (*Notifier)(nil) +) + +type event struct { + kind eventKind + payload string +} + +type Notifier struct { + mu sync.Mutex + cond *sync.Cond + queue *list.List + closed bool + done chan struct{} + + listener listener.NetworkChangeListener + dnsManager dns.IosDnsManager +} + +func New(l listener.NetworkChangeListener, dm dns.IosDnsManager) *Notifier { + n := &Notifier{ + queue: list.New(), + done: make(chan struct{}), + listener: l, + dnsManager: dm, + } + n.cond = sync.NewCond(&n.mu) + go n.deliverLoop() + return n +} + +func (n *Notifier) OnNetworkChanged(routes string) { + n.enqueue(event{kind: eventRoutes, payload: routes}) +} + +func (n *Notifier) SetInterfaceIP(ip string) { + n.enqueue(event{kind: eventIfaceIP, payload: ip}) +} + +func (n *Notifier) SetInterfaceIPv6(ip string) { + n.enqueue(event{kind: eventIfaceIPv6, payload: ip}) +} + +func (n *Notifier) ApplyDns(config string) { + n.enqueue(event{kind: eventDNS, payload: config}) +} + +// Close stops accepting new events and blocks until the delivery loop has +// drained all queued events and exited. +func (n *Notifier) Close() { + n.mu.Lock() + n.closed = true + n.cond.Signal() + n.mu.Unlock() + <-n.done +} + +func (n *Notifier) enqueue(ev event) { + n.mu.Lock() + defer n.mu.Unlock() + if n.closed { + return + } + n.queue.PushBack(ev) + n.cond.Signal() +} + +func (n *Notifier) deliverLoop() { + defer close(n.done) + for { + n.mu.Lock() + for n.queue.Len() == 0 && !n.closed { + n.cond.Wait() + } + if n.closed && n.queue.Len() == 0 { + n.mu.Unlock() + return + } + ev := n.queue.Remove(n.queue.Front()).(event) + l := n.listener + dm := n.dnsManager + n.mu.Unlock() + + switch ev.kind { + case eventRoutes: + if l != nil { + l.OnNetworkChanged(ev.payload) + } + case eventIfaceIP: + if l != nil { + l.SetInterfaceIP(ev.payload) + } + case eventIfaceIPv6: + if l != nil { + l.SetInterfaceIPv6(ev.payload) + } + case eventDNS: + if dm != nil { + dm.ApplyDns(ev.payload) + } + } + } +} diff --git a/client/internal/tunnelnotifier/notifier_test.go b/client/internal/tunnelnotifier/notifier_test.go new file mode 100644 index 000000000..ffbcdc15c --- /dev/null +++ b/client/internal/tunnelnotifier/notifier_test.go @@ -0,0 +1,192 @@ +package tunnelnotifier + +import ( + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type call struct { + kind string + payload string +} + +type recorder struct { + mu sync.Mutex + calls []call + inFlight atomic.Int32 + overlap atomic.Bool + delay time.Duration +} + +func (r *recorder) record(kind, payload string) { + if r.inFlight.Add(1) != 1 { + r.overlap.Store(true) + } + if r.delay > 0 { + time.Sleep(r.delay) + } + r.mu.Lock() + r.calls = append(r.calls, call{kind: kind, payload: payload}) + r.mu.Unlock() + r.inFlight.Add(-1) +} + +func (r *recorder) count() int { + r.mu.Lock() + defer r.mu.Unlock() + return len(r.calls) +} + +func (r *recorder) snapshot() []call { + r.mu.Lock() + defer r.mu.Unlock() + out := make([]call, len(r.calls)) + copy(out, r.calls) + return out +} + +type fakeListener struct { + rec *recorder +} + +func (f *fakeListener) OnNetworkChanged(routes string) { + f.rec.record("routes", routes) +} + +func (f *fakeListener) SetInterfaceIP(ip string) { + f.rec.record("ip", ip) +} + +func (f *fakeListener) SetInterfaceIPv6(ip string) { + f.rec.record("ipv6", ip) +} + +type fakeDNSManager struct { + rec *recorder +} + +func (f *fakeDNSManager) ApplyDns(config string) { + f.rec.record("dns", config) +} + +func TestFIFOOrder(t *testing.T) { + rec := &recorder{} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + defer n.Close() + + n.SetInterfaceIP("10.0.0.1") + n.SetInterfaceIPv6("fd00::1") + n.ApplyDns(`{"domains":[]}`) + n.OnNetworkChanged("10.0.0.0/8,192.168.0.0/16") + n.ApplyDns(`{"domains":["example.com"]}`) + + require.Eventually(t, func() bool { return rec.count() == 5 }, time.Second, time.Millisecond) + + expected := []call{ + {kind: "ip", payload: "10.0.0.1"}, + {kind: "ipv6", payload: "fd00::1"}, + {kind: "dns", payload: `{"domains":[]}`}, + {kind: "routes", payload: "10.0.0.0/8,192.168.0.0/16"}, + {kind: "dns", payload: `{"domains":["example.com"]}`}, + } + assert.Equal(t, expected, rec.snapshot()) +} + +func TestNoOverlappingCalls(t *testing.T) { + rec := &recorder{delay: 100 * time.Microsecond} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + defer n.Close() + + const producers = 8 + const perProducer = 25 + + var wg sync.WaitGroup + for i := 0; i < producers; i++ { + wg.Add(1) + go func(id int) { + defer wg.Done() + for j := 0; j < perProducer; j++ { + payload := fmt.Sprintf("%d-%d", id, j) + switch j % 4 { + case 0: + n.OnNetworkChanged(payload) + case 1: + n.SetInterfaceIP(payload) + case 2: + n.SetInterfaceIPv6(payload) + case 3: + n.ApplyDns(payload) + } + } + }(i) + } + wg.Wait() + + require.Eventually(t, func() bool { return rec.count() == producers*perProducer }, 5*time.Second, time.Millisecond) + assert.False(t, rec.overlap.Load()) +} + +func TestDNSAndRoutesInterleaved(t *testing.T) { + rec := &recorder{delay: 100 * time.Microsecond} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + defer n.Close() + + const events = 50 + + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for i := 0; i < events; i++ { + n.ApplyDns(fmt.Sprintf("dns-%d", i)) + } + }() + go func() { + defer wg.Done() + for i := 0; i < events; i++ { + n.OnNetworkChanged(fmt.Sprintf("routes-%d", i)) + } + }() + wg.Wait() + + require.Eventually(t, func() bool { return rec.count() == 2*events }, 5*time.Second, time.Millisecond) + assert.False(t, rec.overlap.Load()) + + var dnsSeen, routesSeen int + for _, c := range rec.snapshot() { + switch c.kind { + case "dns": + assert.Equal(t, fmt.Sprintf("dns-%d", dnsSeen), c.payload) + dnsSeen++ + case "routes": + assert.Equal(t, fmt.Sprintf("routes-%d", routesSeen), c.payload) + routesSeen++ + } + } + assert.Equal(t, events, dnsSeen) + assert.Equal(t, events, routesSeen) +} + +func TestCloseDrainsQueue(t *testing.T) { + rec := &recorder{delay: time.Millisecond} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + + const events = 20 + for i := 0; i < events; i++ { + n.OnNetworkChanged(fmt.Sprintf("routes-%d", i)) + } + n.Close() + + require.Equal(t, events, rec.count(), "Close must not return before all queued events are delivered") + + n.OnNetworkChanged("after-close") + n.ApplyDns("after-close") + time.Sleep(50 * time.Millisecond) + assert.Equal(t, events, rec.count()) +} diff --git a/client/internal/updater/installer/doc.go b/client/internal/updater/installer/doc.go index 0a60454bb..11b0512ac 100644 --- a/client/internal/updater/installer/doc.go +++ b/client/internal/updater/installer/doc.go @@ -37,23 +37,32 @@ // Updater Process (Setup): // // 1. Receives parameters from service via command-line arguments -// 2. Runs installer with appropriate silent/quiet flags: +// 2. Terminates the UI so the installer does not have to replace a locked image +// file, which would otherwise leave the install needing a reboot +// 3. Runs installer with appropriate silent/quiet flags: // - Windows EXE: installer.exe /S -// - Windows MSI: msiexec.exe /i installer.msi /quiet /qn /l*v msi.log +// - Windows MSI: msiexec.exe /i installer.msi /qn /norestart REBOOT=ReallySuppress /l*v msi.log // - macOS PKG: installer -pkg installer.pkg -target / // - macOS Homebrew: brew upgrade netbirdio/tap/netbird -// 3. Installer terminates daemon and UI processes -// 4. Installer replaces binaries with new version -// 5. Updater waits for installer to complete -// 6. Updater restarts daemon: +// 4. Installer terminates the daemon +// 5. Installer replaces binaries with new version +// 6. Updater waits for installer to complete. On Windows, MSI exit codes 3010 +// (ERROR_SUCCESS_REBOOT_REQUIRED) and 1641 (ERROR_SUCCESS_REBOOT_INITIATED) +// are a pending-reboot outcome, not a failure: the install succeeded, but +// some files are only replaced on the next restart (the reboot itself is +// suppressed via /norestart and REBOOT=ReallySuppress), and the flow +// continues as on success +// 7. Updater restarts daemon: // - Windows: netbird.exe service start // - macOS/Linux: netbird service start -// 7. Updater restarts UI: -// - Windows: Launches netbird-ui.exe as active console user using CreateProcessAsUser +// 8. Updater restarts UI: +// - Windows: Launches netbird-ui.exe using CreateProcessAsUser in every +// session it was terminated in, falling back to the active console session // - macOS: Uses launchctl asuser to launch NetBird.app for console user // - Linux: Not implemented (UI typically auto-starts) -// 8. Updater writes result.json with success/error status -// 9. Updater process exits +// 9. Updater writes result.json with success/error status (a pending reboot is +// recorded as success) +// 10. Updater process exits // // # Result Communication // diff --git a/client/internal/updater/installer/installer_common.go b/client/internal/updater/installer/installer_common.go index 8e44bee82..17566f7de 100644 --- a/client/internal/updater/installer/installer_common.go +++ b/client/internal/updater/installer/installer_common.go @@ -42,6 +42,9 @@ func NewWithDir(tempDir string) *Installer { // This will run by the original service process func (u *Installer) RunInstallation(ctx context.Context, targetVersion string) (err error) { resultHandler := NewResultHandler(u.tempDir) + if err := resultHandler.ClearStaleResult(); err != nil { + log.Warnf("clear stale installer result: %v", err) + } defer func() { if err != nil { diff --git a/client/internal/updater/installer/installer_run_darwin.go b/client/internal/updater/installer/installer_run_darwin.go index 248a404aa..5650bc769 100644 --- a/client/internal/updater/installer/installer_run_darwin.go +++ b/client/internal/updater/installer/installer_run_darwin.go @@ -98,47 +98,44 @@ func (u *Installer) startDaemon(daemonFolder string) error { func (u *Installer) startUIAsUser() error { log.Infof("starting netbird-ui: %s", uiBinary) - // Get the current console user - cmd := exec.Command("stat", "-f", "%Su", "/dev/console") - output, err := cmd.Output() + username, err := consoleUser() if err != nil { - return fmt.Errorf("failed to get console user: %w", err) + return err } - username := strings.TrimSpace(string(output)) - if username == "" || username == "root" { - return fmt.Errorf("no active user session found") - } - - log.Infof("starting UI for user: %s", username) - - // Get user's UID userInfo, err := user.Lookup(username) if err != nil { - return fmt.Errorf("failed to lookup user %s: %w", username, err) + return fmt.Errorf("lookup user %s: %w", username, err) } - // Start the UI process as the console user using launchctl - // This ensures the app runs in the user's context with proper GUI access - launchCmd := exec.Command("launchctl", "asuser", userInfo.Uid, "open", "-a", uiBinary) + log.Infof("starting UI for user: %s (uid %s)", username, userInfo.Uid) + + launchCmd := exec.Command("launchctl", "asuser", userInfo.Uid, "sudo", "-u", username, "-H", "open", "-a", uiBinary) log.Infof("launchCmd: %s", launchCmd.String()) - // Set the user's home directory for proper macOS app behavior - launchCmd.Env = append(os.Environ(), "HOME="+userInfo.HomeDir) - log.Infof("set HOME environment variable: %s", userInfo.HomeDir) - if err := launchCmd.Start(); err != nil { - return fmt.Errorf("failed to start UI process: %w", err) - } - - // Release the process so it can run independently - if err := launchCmd.Process.Release(); err != nil { - log.Warnf("failed to release UI process: %v", err) + if err := launchCmd.Run(); err != nil { + return fmt.Errorf("run UI launch: %w", err) } log.Infof("netbird-ui started successfully for user %s", username) return nil } +func consoleUser() (string, error) { + output, err := exec.Command("stat", "-f", "%Su", "/dev/console").Output() + if err != nil { + return "", fmt.Errorf("get console user: %w", err) + } + + username := strings.TrimSpace(string(output)) + switch username { + case "", "root", "loginwindow", "_mbsetupuser": + return "", fmt.Errorf("no active GUI user session, console user: %q", username) + } + + return username, nil +} + func (u *Installer) installPkgFile(ctx context.Context, path string) error { log.Infof("installing pkg file: %s", path) diff --git a/client/internal/updater/installer/installer_run_windows.go b/client/internal/updater/installer/installer_run_windows.go index 70c7e32cf..b2ecf3299 100644 --- a/client/internal/updater/installer/installer_run_windows.go +++ b/client/internal/updater/installer/installer_run_windows.go @@ -2,6 +2,7 @@ package installer import ( "context" + "errors" "fmt" "os" "os/exec" @@ -22,6 +23,12 @@ const ( msiLogFile = "msi.log" + // ERROR_SUCCESS_REBOOT_REQUIRED and ERROR_SUCCESS_REBOOT_INITIATED + msiRebootRequired = 3010 + msiRebootInitiated = 1641 + + processExitWait = 10 * time.Second + msiDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.msi" exeDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.exe" ) @@ -38,6 +45,8 @@ var ( func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string, daemonFolder string) (resultErr error) { resultHandler := NewResultHandler(u.tempDir) + var uiSessions []uint32 + // Always ensure daemon and UI are restarted after setup defer func() { log.Infof("starting daemon back") @@ -46,7 +55,7 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string } log.Infof("starting UI back") - if err := u.startUIAsUser(daemonFolder); err != nil { + if err := u.startUI(daemonFolder, uiSessions); err != nil { log.Errorf("failed to start UI: %v", err) } @@ -75,6 +84,14 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string return } + // The UI holds an open handle on its own image. Left running, Restart Manager + // cannot shut it down (msiexec runs as LocalSystem here, the UI as the + // interactive user), so the MSI falls back to replacing the file on reboot and + // marks the install as restart-required. The deferred close-application action + // in the package runs too late to prevent that, it happens after + // InstallValidate has already registered the file as in use. + uiSessions = killUI() + var cmd *exec.Cmd switch installerType { case TypeExe: @@ -84,7 +101,9 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string installerDir := filepath.Dir(installerFile) logPath := filepath.Join(installerDir, msiLogFile) log.Infof("run msi installer: %s", installerFile) - cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/quiet", "/qn", "/l*v", logPath) + // REBOOT=ReallySuppress: a silent install has no way to ask, so without it + // msiexec reboots the machine on its own if it decides one is needed. + cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/qn", "/norestart", "REBOOT=ReallySuppress", "/l*v", logPath) } cmd.Dir = filepath.Dir(installerFile) @@ -95,9 +114,13 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string } log.Infof("installer started with PID %d", cmd.Process.Pid) - if resultErr = cmd.Wait(); resultErr != nil { - log.Errorf("installer process finished with error: %v", resultErr) - return + if err := cmd.Wait(); err != nil { + if !isRebootPending(err) { + resultErr = err + log.Errorf("installer process finished with error: %v", err) + return + } + log.Warnf("installer completed but reported a pending reboot, some files will be replaced on the next restart") } return nil @@ -117,16 +140,142 @@ func (u *Installer) startDaemon(daemonFolder string) error { return nil } -func (u *Installer) startUIAsUser(daemonFolder string) error { +func (u *Installer) startUI(daemonFolder string, sessionIDs []uint32) error { uiPath := filepath.Join(daemonFolder, uiName) log.Infof("starting netbird-ui: %s", uiPath) - // Get the active console session ID - sessionID := windows.WTSGetActiveConsoleSessionId() - if sessionID == 0xFFFFFFFF { - return fmt.Errorf("no active user session found") + if len(sessionIDs) == 0 { + sessionID := windows.WTSGetActiveConsoleSessionId() + if sessionID == 0xFFFFFFFF { + return fmt.Errorf("no active user session found") + } + sessionIDs = []uint32{sessionID} } + var errs []error + for _, sessionID := range sessionIDs { + if err := startUIInSession(uiPath, sessionID); err != nil { + errs = append(errs, fmt.Errorf("session %d: %w", sessionID, err)) + continue + } + log.Infof("netbird-ui started successfully in session %d", sessionID) + } + return errors.Join(errs...) +} + +// isRebootPending reports whether the installer exit code means it succeeded but +// left work for the next restart. The reboot itself is suppressed, so this is not +// a failure. +func isRebootPending(err error) bool { + var exitErr *exec.ExitError + if !errors.As(err, &exitErr) { + return false + } + + switch exitErr.ExitCode() { + case msiRebootRequired, msiRebootInitiated: + return true + default: + return false + } +} + +// killUI terminates any running netbird-ui process and returns the IDs of the +// interactive sessions the terminated processes belonged to. Setup starts the +// UI again in those sessions once the installer is done. +func killUI() []uint32 { + pids, err := processIDsByName(uiName) + if err != nil { + log.Warnf("failed to look up %s processes: %v", uiName, err) + return nil + } + + sessions := make(map[uint32]struct{}) + for _, pid := range pids { + var sessionID uint32 + if err := windows.ProcessIdToSessionId(pid, &sessionID); err != nil { + log.Warnf("failed to look up session of %s (PID %d): %v", uiName, pid, err) + } + + if err := terminateProcess(pid); err != nil { + log.Warnf("failed to terminate %s (PID %d): %v", uiName, pid, err) + continue + } + log.Infof("terminated %s (PID %d) in session %d", uiName, pid, sessionID) + + if sessionID != 0 { + sessions[sessionID] = struct{}{} + } + } + + sessionIDs := make([]uint32, 0, len(sessions)) + for sessionID := range sessions { + sessionIDs = append(sessionIDs, sessionID) + } + return sessionIDs +} + +func processIDsByName(name string) ([]uint32, error) { + snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0) + if err != nil { + return nil, fmt.Errorf("create process snapshot: %w", err) + } + defer func() { + if err := windows.CloseHandle(snapshot); err != nil { + log.Warnf("failed to close process snapshot: %v", err) + } + }() + + var entry windows.ProcessEntry32 + entry.Size = uint32(unsafe.Sizeof(entry)) + + var pids []uint32 + for err = windows.Process32First(snapshot, &entry); err == nil; err = windows.Process32Next(snapshot, &entry) { + if strings.EqualFold(windows.UTF16ToString(entry.ExeFile[:]), name) { + pids = append(pids, entry.ProcessID) + } + } + if !errors.Is(err, windows.ERROR_NO_MORE_FILES) { + return nil, fmt.Errorf("enumerate processes: %w", err) + } + + return pids, nil +} + +func terminateProcess(pid uint32) error { + handle, err := windows.OpenProcess(windows.PROCESS_TERMINATE|windows.SYNCHRONIZE, false, pid) + if err != nil { + // The process may have exited between enumeration and now. + if errors.Is(err, windows.ERROR_INVALID_PARAMETER) { + return nil + } + return fmt.Errorf("open process: %w", err) + } + defer func() { + if err := windows.CloseHandle(handle); err != nil { + log.Warnf("failed to close process handle: %v", err) + } + }() + + if err := windows.TerminateProcess(handle, 0); err != nil { + return fmt.Errorf("terminate process: %w", err) + } + + // Wait for the handle to signal so the image file is released before the + // installer tries to overwrite it. A timeout is reported through the returned + // event, not through err, which stays nil unless the wait itself failed. + event, err := windows.WaitForSingleObject(handle, uint32(processExitWait.Milliseconds())) + if err != nil { + return fmt.Errorf("wait for process exit: %w", err) + } + if event != windows.WAIT_OBJECT_0 { + return fmt.Errorf("wait for process exit: unexpected wait result %#x", event) + } + + return nil +} + +func startUIInSession(uiPath string, sessionID uint32) error { // Get the user token for that session var userToken windows.Token err := windows.WTSQueryUserToken(sessionID, &userToken) @@ -158,6 +307,16 @@ func (u *Installer) startUIAsUser(daemonFolder string) error { } }() + var env *uint16 + if err := windows.CreateEnvironmentBlock(&env, primaryToken, false); err != nil { + return fmt.Errorf("create environment block: %w", err) + } + defer func() { + if err := windows.DestroyEnvironmentBlock(env); err != nil { + log.Warnf("failed to destroy environment block: %v", err) + } + }() + // Prepare startup info var si windows.StartupInfo si.Cb = uint32(unsafe.Sizeof(si)) @@ -180,7 +339,7 @@ func (u *Installer) startUIAsUser(daemonFolder string) error { nil, false, creationFlags, - nil, + env, nil, &si, &pi, @@ -197,7 +356,6 @@ func (u *Installer) startUIAsUser(daemonFolder string) error { log.Warnf("failed to close thread handle: %v", err) } - log.Infof("netbird-ui started successfully in session %d", sessionID) return nil } diff --git a/client/internal/updater/installer/installer_run_windows_test.go b/client/internal/updater/installer/installer_run_windows_test.go new file mode 100644 index 000000000..6a4540610 --- /dev/null +++ b/client/internal/updater/installer/installer_run_windows_test.go @@ -0,0 +1,108 @@ +package installer + +import ( + "errors" + "os/exec" + "slices" + "strconv" + "testing" +) + +// exitErrorWithCode returns a real *exec.ExitError carrying the given exit code. +func exitErrorWithCode(t *testing.T, code int) error { + t.Helper() + + err := exec.Command("cmd.exe", "/c", "exit "+strconv.Itoa(code)).Run() + if err == nil { + t.Fatalf("expected a non-zero exit for code %d", code) + } + return err +} + +func TestIsRebootPending(t *testing.T) { + tests := []struct { + name string + code int + want bool + }{ + {name: "reboot required", code: msiRebootRequired, want: true}, + {name: "reboot initiated", code: msiRebootInitiated, want: true}, + {name: "generic failure", code: 1603, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isRebootPending(exitErrorWithCode(t, tt.code)); got != tt.want { + t.Errorf("isRebootPending(exit %d) = %v, want %v", tt.code, got, tt.want) + } + }) + } +} + +// TestProcessIDsByNameAndTerminate spawns a long-running system process, finds it +// by name and terminates it, covering the path the updater uses to release the UI +// image file before the installer replaces it. +func TestProcessIDsByNameAndTerminate(t *testing.T) { + cmd := exec.Command("ping.exe", "-n", "60", "127.0.0.1") + if err := cmd.Start(); err != nil { + t.Fatalf("start ping: %v", err) + } + + pid := uint32(cmd.Process.Pid) + killed := false + t.Cleanup(func() { + if !killed { + _ = cmd.Process.Kill() + } + _ = cmd.Wait() + }) + + // Name matching must be case-insensitive: the snapshot reports PING.EXE. + pids, err := processIDsByName("ping.exe") + if err != nil { + t.Fatalf("processIDsByName: %v", err) + } + + if !slices.Contains(pids, pid) { + t.Fatalf("PID %d not among the ping.exe processes found: %v", pid, pids) + } + + if err := terminateProcess(pid); err != nil { + t.Fatalf("terminateProcess: %v", err) + } + killed = true + + // terminateProcess only returns once the handle has signalled, so the process + // is already gone and Wait must not block. It exits with the code passed to + // TerminateProcess, which is 0, so Wait reports no error. + if err := cmd.Wait(); err != nil { + t.Fatalf("wait for terminated ping: %v", err) + } + if !cmd.ProcessState.Exited() { + t.Error("process did not exit after terminateProcess") + } + + remaining, err := processIDsByName("ping.exe") + if err != nil { + t.Fatalf("processIDsByName after terminate: %v", err) + } + if slices.Contains(remaining, pid) { + t.Errorf("PID %d still listed after terminateProcess", pid) + } +} + +func TestProcessIDsByNameNoMatch(t *testing.T) { + pids, err := processIDsByName("netbird-nonexistent-process.exe") + if err != nil { + t.Fatalf("processIDsByName: %v", err) + } + if len(pids) != 0 { + t.Errorf("expected no matches, got %v", pids) + } +} + +func TestIsRebootPendingNonExitError(t *testing.T) { + if isRebootPending(errors.New("start installer: file not found")) { + t.Error("a non-exit error must not be treated as a pending reboot") + } +} diff --git a/client/internal/updater/installer/result.go b/client/internal/updater/installer/result.go index 526c3eb53..55a0d8ac8 100644 --- a/client/internal/updater/installer/result.go +++ b/client/internal/updater/installer/result.go @@ -54,6 +54,12 @@ func (rh *ResultHandler) GetErrorResultReason() string { return "" } +// ClearStaleResult removes a result file left over from a previous installation +// attempt so result watchers cannot read an outdated outcome for the current attempt. +func (rh *ResultHandler) ClearStaleResult() error { + return rh.cleanup() +} + func (rh *ResultHandler) WriteSuccess() error { result := Result{ Success: true, diff --git a/client/internal/updater/manager.go b/client/internal/updater/manager.go index 7fc300739..1b69368d0 100644 --- a/client/internal/updater/manager.go +++ b/client/internal/updater/manager.go @@ -435,7 +435,7 @@ func (m *Manager) install(ctx context.Context, pendingVersion *v.Version) error } inst := installer.New() - if err := inst.RunInstallation(ctx, pendingVersion.String()); err != nil { + if err := inst.RunInstallation(ctx, pendingVersion.String()); err != nil { //nolint:staticcheck // always errors on platforms without an installer log.Errorf("error triggering update: %v", err) m.statusRecorder.PublishEvent( cProto.SystemEvent_ERROR, diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index a2f123900..8373e498a 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -13,8 +13,8 @@ import ( "time" log "github.com/sirupsen/logrus" - "golang.org/x/exp/maps" + nbAnonymize "github.com/netbirdio/netbird/client/anonymize" "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/debug" @@ -22,6 +22,7 @@ import ( "github.com/netbirdio/netbird/client/internal/listener" "github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/netevents" "github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/formatter" "github.com/netbirdio/netbird/route" @@ -29,10 +30,12 @@ import ( types "github.com/netbirdio/netbird/upload-server/types" ) -// ConnectionListener export internal Listener for mobile -type ConnectionListener interface { - peer.Listener -} +// AnonymizeLevelDefault and AnonymizeLevelStrict are the accepted +// anonymizeLevel values for DebugBundle. +const ( + AnonymizeLevelDefault = nbAnonymize.LevelDefaultString + AnonymizeLevelStrict = nbAnonymize.LevelStrictString +) // RouteListener export internal RouteListener for mobile type NetworkChangeListener interface { @@ -80,6 +83,10 @@ type Client struct { onHostDnsFn func([]string) dnsManager dns.IosDnsManager loginComplete bool + // netMgr outlives engine restarts: it mirrors the OS connectivity, not + // the engine lifecycle. Run injects its state and sweeper into each new + // ConnectClient. + netMgr *netevents.Manager // preloadedConfig holds config loaded from JSON (used on tvOS where file writes are blocked) preloadedConfig *profilemanager.Config @@ -90,6 +97,7 @@ type Client struct { // NewClient instantiate a new Client func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osVersion string, osName string, networkChangeListener NetworkChangeListener, dnsManager DnsManager) *Client { + recorder := peer.NewRecorder("") return &Client{ cfgFile: cfgFile, stateFile: stateFile, @@ -98,10 +106,11 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV deviceName: deviceName, osName: osName, osVersion: osVersion, - recorder: peer.NewRecorder(""), + recorder: recorder, ctxCancelLock: &sync.Mutex{}, networkChangeListener: networkChangeListener, dnsManager: dnsManager, + netMgr: netevents.NewManager(recorder), } } @@ -159,19 +168,26 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error { defer c.ctxCancel() c.ctxCancelLock.Unlock() - auth := NewAuthWithConfig(ctx, cfg) - err = auth.LoginSync() - if err != nil { - return err - } - - log.Infof("Auth successful") + // No login pre-flight here. The engine's own loginToManagement (connect.go) performs + // the authoritative Login immediately before the first Sync, so a LoginSync() call at + // this point only duplicated it — costing two extra Login RPCs (IsLoginRequired + + // Login) on every engine start, since IsLoginRequired is itself a full Login RPC. + // + // Auth failures still reach the caller through the engine path: loginToManagement + // returns PermissionDenied, which marks the shared status recorder + // (MarkManagementDisconnected) and fires ClientStop → onDisconnected, where + // IsLoginRequiredCached() reports login-required. The error is also returned out of Run(). + // + // A pre-flight was also actively harmful when the server is unreachable: its 2-minute + // backoff blocked the start and then reported "login required" for what was really a + // timeout. The engine instead keeps retrying and recovers when the server returns. // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) c.onHostDnsFn = func([]string) {} cfg.WgIface = interfaceName - connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) + connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, + internal.WithNetEvents(c.netMgr)) c.setState(cfg, connectClient) // Persist the latest sync response so DebugBundle can include the network // map. On iOS this is backed by disk to keep it out of the constrained @@ -180,6 +196,25 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error { return connectClient.RunOniOS(fd, c.networkChangeListener, c.dnsManager, c.stateFile, c.cacheDir, c.logFilePath) } +// SetNetworkAvailable feeds OS-reported network availability into the client +// (e.g. from NWPathMonitor). While unavailable, the internal reconnect loops +// suspend their attempts and the connection listener reports NoNetwork +// instead of Connecting; when availability returns, the loops resume +// immediately with a fresh backoff. Losing the last network also sweeps the +// registered connections, so the client does not keep reporting Connected +// over stale sockets with no network at all. +func (c *Client) SetNetworkAvailable(available bool) { + c.netMgr.SetNetworkAvailable(available) +} + +// NotifyNetworkChange marks the management, signal and relay connections +// stale after the OS switched networks and schedules a sweep that cuts +// whatever has not redialed on the new network by then. The engine and the +// TUN device stay untouched. +func (c *Client) NotifyNetworkChange() { + c.netMgr.NotifyNetworkChange() +} + // Stop the internal client and free the resources func (c *Client) Stop() { c.ctxCancelLock.Lock() @@ -195,8 +230,10 @@ func (c *Client) Stop() { // DebugBundle generates a debug bundle, uploads it and returns the upload key. // It works with or without a running engine: when the engine is up it reuses // the live config, sync response and client metrics; otherwise it loads the -// config from disk (or the preloaded tvOS config). -func (c *Client) DebugBundle(anonymize bool) (string, error) { +// config from disk (or the preloaded tvOS config). anonymizeLevel is "default" +// or "strict"; strict also anonymizes internal IP ranges, peer names, and +// WireGuard public keys, and implies anonymize. +func (c *Client) DebugBundle(anonymize bool, anonymizeLevel string) (string, error) { cfg, cc := c.stateSnapshot() // If the engine hasn't been started, load config so we can reach management. @@ -246,6 +283,7 @@ func (c *Client) DebugBundle(anonymize bool) (string, error) { deps, debug.BundleConfig{ Anonymize: anonymize, + AnonymizeLevel: nbAnonymize.ParseLevel(anonymizeLevel), IncludeSystemInfo: true, }, ) @@ -263,7 +301,7 @@ func (c *Client) DebugBundle(anonymize bool) (string, error) { uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() - key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path) + key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path, false) if err != nil { return "", fmt.Errorf("upload debug bundle: %w", err) } @@ -315,7 +353,11 @@ func (c *Client) GetStatusDetails() *StatusDetails { // SetConnectionListener set the network connection listener func (c *Client) SetConnectionListener(listener ConnectionListener) { - c.recorder.SetConnectionListener(listener) + if listener == nil { + c.recorder.RemoveConnectionListener() + return + } + c.recorder.SetConnectionListener(connectionListenerAdapter{listener}) } // RemoveConnectionListener remove connection listener @@ -637,23 +679,18 @@ func (c *Client) SelectRoute(id string) error { } routeManager := engine.GetRouteManager() - routeSelector := routeManager.GetRouteSelector() if id == "All" { log.Debugf("select all routes") - routeSelector.SelectAllRoutes() - } else { - log.Debugf("select route with id: %s", id) - routes := toNetIDs([]string{id}) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - if err := routeSelector.SelectRoutes(routes, true, maps.Keys(routesMap)); err != nil { - log.Debugf("error when selecting routes: %s", err) - return fmt.Errorf("select routes: %w", err) - } + routeManager.SelectAllRoutes() + return nil } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) - return nil + log.Debugf("select route with id: %s", id) + if err := routeManager.SelectRoutes(toNetIDs([]string{id}), true); err != nil { + log.Debugf("error when selecting routes: %s", err) + return err + } + return nil } func (c *Client) DeselectRoute(id string) error { @@ -667,21 +704,17 @@ func (c *Client) DeselectRoute(id string) error { } routeManager := engine.GetRouteManager() - routeSelector := routeManager.GetRouteSelector() if id == "All" { log.Debugf("deselect all routes") - routeSelector.DeselectAllRoutes() - } else { - log.Debugf("deselect route with id: %s", id) - routes := toNetIDs([]string{id}) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - if err := routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil { - log.Debugf("error when deselecting routes: %s", err) - return fmt.Errorf("deselect routes: %w", err) - } + routeManager.DeselectAllRoutes() + return nil + } + + log.Debugf("deselect route with id: %s", id) + if err := routeManager.DeselectRoutes(toNetIDs([]string{id})); err != nil { + log.Debugf("error when deselecting routes: %s", err) + return err } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) return nil } diff --git a/client/ios/NetBirdSDK/connection_listener.go b/client/ios/NetBirdSDK/connection_listener.go new file mode 100644 index 000000000..d792537ba --- /dev/null +++ b/client/ios/NetBirdSDK/connection_listener.go @@ -0,0 +1,43 @@ +//go:build ios + +package NetBirdSDK + +import ( + "github.com/netbirdio/netbird/client/internal/peer" +) + +// Client state values, re-exported as basic constants so gomobile emits them +// into the generated bindings. They mirror peer.ClientState*: append-only, +// never reorder. +const ( + ClientStateDisconnected = int(peer.ClientStateDisconnected) + ClientStateConnected = int(peer.ClientStateConnected) + ClientStateConnecting = int(peer.ClientStateConnecting) + ClientStateDisconnecting = int(peer.ClientStateDisconnecting) + ClientStateNoNetwork = int(peer.ClientStateNoNetwork) +) + +// ConnectionListener export internal Listener for mobile. +// +// It intentionally lacks OnStateChanged for now: adding a method to a gomobile +// interface breaks every Swift implementation, so the iOS app keeps building +// against the legacy per-state callbacks. A follow-up will extend it together +// with the app. +type ConnectionListener interface { + OnConnected() + OnDisconnected() + OnConnecting() + OnDisconnecting() + OnAddressChanged(string, string) + OnPeersListChanged(int) +} + +// connectionListenerAdapter adapts the gomobile-facing ConnectionListener to +// peer.Listener. +type connectionListenerAdapter struct { + ConnectionListener +} + +// OnStateChanged is dropped on iOS until the app adopts the state callback; +// the legacy per-state callbacks continue to fire. +func (a connectionListenerAdapter) OnStateChanged(peer.ClientState) {} diff --git a/client/ios/NetBirdSDK/login.go b/client/ios/NetBirdSDK/login.go index 99486839b..42a575359 100644 --- a/client/ios/NetBirdSDK/login.go +++ b/client/ios/NetBirdSDK/login.go @@ -222,17 +222,36 @@ func (a *Auth) Login(resultListener ErrListener, urlOpener URLOpener, forceDevic // LoginWithDeviceName performs interactive login with device authentication support // The deviceName parameter allows specifying a custom device name (required for tvOS) func (a *Auth) LoginWithDeviceName(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string) { + a.startLogin(resultListener, urlOpener, forceDeviceAuth, deviceName, false) +} + +// LoginInteractive performs the same interactive login as LoginWithDeviceName but skips the +// IsLoginRequired() pre-flight and goes straight to the browser / device-code flow. +// +// IsLoginRequired() is itself a full Login RPC against the management server, so when the +// caller has ALREADY established that login is required it is a pure duplicate. On iOS the +// main app decides to show the browser based on its own isLoginRequired() check and then +// calls straight into this method, so re-asking the server would add another Login RPC to +// every interactive login. +// +// Use LoginWithDeviceName when the auth state is unknown and a silent (browser-less) login +// must still be possible; use this when the browser is going to be shown regardless. +func (a *Auth) LoginInteractive(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string) { + a.startLogin(resultListener, urlOpener, forceDeviceAuth, deviceName, true) +} + +func (a *Auth) startLogin(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string, skipLoginCheck bool) { if resultListener == nil { - log.Errorf("LoginWithDeviceName: resultListener is nil") + log.Errorf("startLogin: resultListener is nil") return } if urlOpener == nil { - log.Errorf("LoginWithDeviceName: urlOpener is nil") + log.Errorf("startLogin: urlOpener is nil") resultListener.OnError(fmt.Errorf("urlOpener is nil")) return } go func() { - err := a.login(urlOpener, forceDeviceAuth, deviceName) + err := a.login(urlOpener, forceDeviceAuth, deviceName, skipLoginCheck) if err != nil { resultListener.OnError(err) } else { @@ -241,7 +260,7 @@ func (a *Auth) LoginWithDeviceName(resultListener ErrListener, urlOpener URLOpen }() } -func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName string) error { +func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName string, skipLoginCheck bool) error { // Create context with device name if provided ctx := a.ctx if deviceName != "" { @@ -255,10 +274,13 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin } defer authClient.Close() - // check if we need to generate JWT token - needsLogin, err := authClient.IsLoginRequired(ctx) - if err != nil { - return fmt.Errorf("failed to check login requirement: %v", err) + // check if we need to generate JWT token (skipped when the caller already knows) + needsLogin := true + if !skipLoginCheck { + needsLogin, err = authClient.IsLoginRequired(ctx) + if err != nil { + return fmt.Errorf("failed to check login requirement: %v", err) + } } jwtToken := "" @@ -301,7 +323,7 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin const authInfoRequestTimeout = 30 * time.Second func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) { - oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth) + oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth, "") if err != nil { return nil, fmt.Errorf("failed to get OAuth flow: %v", err) } diff --git a/client/ios/NetBirdSDK/version.go b/client/ios/NetBirdSDK/version.go new file mode 100644 index 000000000..606ad18e2 --- /dev/null +++ b/client/ios/NetBirdSDK/version.go @@ -0,0 +1,12 @@ +//go:build ios + +package NetBirdSDK + +import "github.com/netbirdio/netbird/version" + +// GoClientVersion returns the NetBird Go client version that was baked into +// the framework at compile time via +// -ldflags "-X github.com/netbirdio/netbird/version.version=". +func GoClientVersion() string { + return version.NetbirdVersion() +} diff --git a/client/jobexec/executor.go b/client/jobexec/executor.go index e29cc8840..9401acacc 100644 --- a/client/jobexec/executor.go +++ b/client/jobexec/executor.go @@ -54,7 +54,7 @@ func (e *Executor) BundleJob(ctx context.Context, debugBundleDependencies debug. } }() - key, err := debug.UploadDebugBundle(ctx, types.DefaultBundleURL, mgmURL, path) + key, err := debug.UploadDebugBundle(ctx, types.DefaultBundleURL, mgmURL, path, false) if err != nil { log.Errorf("failed to upload debug bundle: %v", err) return "", fmt.Errorf("upload debug bundle: %w", err) diff --git a/client/netevents/netevents.go b/client/netevents/netevents.go new file mode 100644 index 000000000..474cbfa22 --- /dev/null +++ b/client/netevents/netevents.go @@ -0,0 +1,173 @@ +// Package netevents owns the OS network event handling shared by the mobile +// bindings: availability changes park or wake the reconnection loops and drive +// the NoNetwork listener state, and both losing the last network and switching +// networks sweep the stale connections so their owners redial immediately. +package netevents + +import ( + "context" + "sync" + "time" + + "github.com/cenkalti/backoff/v4" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/netevents/netstate" + "github.com/netbirdio/netbird/client/netevents/sweep" +) + +// Recorder receives the availability changes for listener state reporting. +type Recorder interface { + SetNetworkAvailable(available bool) +} + +// Manager ties the network availability state, the connection sweeper and the +// status recorder together; it outlives engine restarts. A nil *Manager is +// the valid no-events value for consumers: the read methods report +// always-online and never sweep. Only the event sources hold a real Manager, +// so the write methods do not tolerate a nil receiver. +type Manager struct { + // mu serializes availability transitions: the IsOnline check and the + // state update must be atomic, or a racing offline flip can skip the sweep + // and leave netState and the recorder disagreeing. + mu sync.Mutex + netState *netstate.State + sweeper *sweep.Sweeper + recorder Recorder +} + +// NewManager creates a Manager reporting into recorder, starting online. +func NewManager(recorder Recorder) *Manager { + return &Manager{ + netState: netstate.New(), + sweeper: sweep.New(), + recorder: recorder, + } +} + +// SetNetworkAvailable records OS-reported network availability. While +// unavailable, the reconnection loops suspend their attempts and the +// connection listener reports NoNetwork instead of Connecting; when +// availability returns, the loops resume immediately with a fresh backoff. +// Losing the last network also sweeps the registered connections: nothing can +// redial while offline, so the stale sockets would otherwise stay silently +// "connected" until their own timeouts and the client would keep reporting +// Connected with no network at all. +// +// Panics on a nil receiver: only the mobile bindings that own a Manager +// report availability. +func (m *Manager) SetNetworkAvailable(available bool) { + m.mu.Lock() + defer m.mu.Unlock() + + if !available && m.netState.IsOnline() { + m.sweeper.MarkNetworkChange() + } + m.netState.Set(available) + m.recorder.SetNetworkAvailable(available) +} + +// NotifyNetworkChange marks the management, signal and relay connections +// stale after the OS switched networks and schedules a sweep that cuts +// whatever has not redialed on the new network by then. The engine and the +// TUN device stay untouched. +// +// Panics on a nil receiver: only the mobile bindings that own a Manager +// report network changes. +func (m *Manager) NotifyNetworkChange() { + m.sweeper.MarkNetworkChange() + log.Infof("network change: connections marked stale") +} + +// IsOnline reports whether the OS reports at least one usable network. +func (m *Manager) IsOnline() bool { + if m == nil { + return true + } + return m.netState.IsOnline() +} + +// Changed returns a channel closed on the next availability transition. +func (m *Manager) Changed() <-chan struct{} { + if m == nil { + return nil + } + return m.netState.Changed() +} + +// Wait blocks while the network is offline; see netstate.State.Wait. +func (m *Manager) Wait(ctx context.Context) (bool, error) { + if m == nil { + return false, nil + } + return m.netState.Wait(ctx) +} + +// WaitSettled waits until an online verdict holds for a full settleWindow, or +// while offline until the budget runs out. Returns false when ctx is +// cancelled. The settle window exists because a disconnect often precedes the +// OS offline flag by a few milliseconds, so a fresh online verdict cannot be +// trusted immediately. A nil Manager has no events to watch: it degrades to a +// fixed budget-long sleep. +func (m *Manager) WaitSettled(ctx context.Context, budget, settleWindow time.Duration) bool { + if m == nil { + select { + case <-time.After(budget): + return true + case <-ctx.Done(): + return false + } + } + + budgetTimer := time.NewTimer(budget) + defer budgetTimer.Stop() + + settle := time.NewTimer(settleWindow) + defer settle.Stop() + + for { + // Channel first, flag second: a flip in between still fires the channel. + changedCh := m.netState.Changed() + if m.netState.IsOnline() { + select { + case <-settle.C: + return true + case <-changedCh: + case <-ctx.Done(): + return false + } + } else { + select { + case <-budgetTimer.C: + return true + case <-changedCh: + case <-ctx.Done(): + return false + } + } + if !settle.Stop() { + select { + case <-settle.C: + default: + } + } + settle.Reset(settleWindow) + } +} + +// StartDial registers an in-flight dial with the sweeper; see sweep.Sweeper.StartDial. +func (m *Manager) StartDial(ctx context.Context) *sweep.Dial { + if m == nil { + return (*sweep.Sweeper)(nil).StartDial(ctx) + } + return m.sweeper.StartDial(ctx) +} + +// QuickRetryBackoff wraps bo for a quick retry after a network change; see +// sweep.Sweeper.QuickRetryBackoff. +func (m *Manager) QuickRetryBackoff(ctx context.Context, bo backoff.BackOff) backoff.BackOff { + if m == nil { + return bo + } + return m.sweeper.QuickRetryBackoff(ctx, bo, m.netState) +} diff --git a/client/netevents/netevents_test.go b/client/netevents/netevents_test.go new file mode 100644 index 000000000..a62ddc270 --- /dev/null +++ b/client/netevents/netevents_test.go @@ -0,0 +1,34 @@ +package netevents + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +type recorderStub struct{} + +func (recorderStub) SetNetworkAvailable(bool) {} + +func TestWaitSettledAfterOutage(t *testing.T) { + const budget = 1500 * time.Millisecond + const settleWindow = 200 * time.Millisecond + const outage = 2 * settleWindow + + m := NewManager(recorderStub{}) + m.SetNetworkAvailable(false) + + start := time.Now() + go func() { + time.Sleep(outage) + m.SetNetworkAvailable(true) + }() + + ok := m.WaitSettled(context.Background(), budget, settleWindow) + elapsed := time.Since(start) + + assert.True(t, ok, "recovered network must let the caller proceed") + assert.GreaterOrEqual(t, elapsed, outage+settleWindow, "an online verdict must hold a full settle window before it is trusted") +} diff --git a/client/netevents/netstate/netstate.go b/client/netevents/netstate/netstate.go new file mode 100644 index 000000000..0d7a1268b --- /dev/null +++ b/client/netevents/netstate/netstate.go @@ -0,0 +1,110 @@ +// Package netstate tracks OS-reported network availability for the client. +// +// A State instance is owned by the platform integration (e.g. the Android or +// iOS bindings, fed from ConnectivityManager callbacks or NWPathMonitor) and +// is injected into the connection retry loops (management, signal, relay, +// peer guards and the top-level connect loop), which consult it to avoid +// burning CPU and battery on reconnect attempts while the device has no +// network at all (e.g. airplane mode), and to reset their backoff as soon as +// the network returns. +// +// Consumers hold a *State that may be nil — every non-mobile platform leaves +// it unset. The read methods are safe on a nil receiver: they report online +// and never block, so consumers behave as if this package did not exist. +package netstate + +import ( + "context" + "sync" + + log "github.com/sirupsen/logrus" +) + +// State holds the OS-reported network availability. The zero value is not +// usable; create instances with New. +type State struct { + mu sync.Mutex + online bool + changed chan struct{} +} + +// New creates a State that starts online. Platforms without network tracking +// pass a nil *State instead: the read methods treat nil as always online and +// never block, so consumers need no nil guards. +func New() *State { + return &State{ + online: true, + changed: make(chan struct{}), + } +} + +// Set records whether the OS reports any usable network. Transitions wake up +// all Wait callers immediately. Unlike the read methods, Set is not nil-safe: +// it is only for the platform owner that created the State with New. +func (s *State) Set(online bool) { + s.mu.Lock() + defer s.mu.Unlock() + if s.online == online { + return + } + s.online = online + close(s.changed) + s.changed = make(chan struct{}) + log.Infof("OS network availability changed: online=%t", online) +} + +// IsOnline reports whether the OS reports at least one usable network. On a +// nil receiver — no State injected — it reports online. +func (s *State) IsOnline() bool { + if s == nil { + return true + } + s.mu.Lock() + defer s.mu.Unlock() + return s.online +} + +// Changed returns a channel closed on the next availability transition, for +// callers that already own a select loop and cannot block in Wait. Re-read it +// after every fire: each transition installs a fresh channel. On a nil +// receiver — no State injected — it returns nil, which blocks forever in a +// select, so the caller simply never observes a transition. +func (s *State) Changed() <-chan struct{} { + if s == nil { + return nil + } + s.mu.Lock() + defer s.mu.Unlock() + return s.changed +} + +// Wait blocks while the network is offline. It reports whether it had to +// wait, so callers can reset their backoff after an outage. It returns early +// with the context error when ctx is done. On a nil receiver — no State +// injected — it returns immediately. +func (s *State) Wait(ctx context.Context) (bool, error) { + if s == nil { + return false, nil + } + waited := false + for { + s.mu.Lock() + if s.online { + s.mu.Unlock() + return waited, nil + } + ch := s.changed + s.mu.Unlock() + + if !waited { + waited = true + log.Debugf("network is offline, pausing connection attempts") + } + + select { + case <-ctx.Done(): + return waited, ctx.Err() + case <-ch: + } + } +} diff --git a/client/netevents/netstate/netstate_test.go b/client/netevents/netstate/netstate_test.go new file mode 100644 index 000000000..ea7015761 --- /dev/null +++ b/client/netevents/netstate/netstate_test.go @@ -0,0 +1,170 @@ +package netstate + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestNewStateIsOnline(t *testing.T) { + assert.True(t, New().IsOnline(), "a fresh State should start online") +} + +func TestSetTogglesOnlineState(t *testing.T) { + s := New() + + s.Set(false) + assert.False(t, s.IsOnline(), "state should be offline after Set(false)") + + s.Set(true) + assert.True(t, s.IsOnline(), "state should be online after Set(true)") +} + +func TestWaitReturnsImmediatelyWhenOnline(t *testing.T) { + s := New() + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + waited, err := s.Wait(ctx) + require.NoError(t, err) + assert.False(t, waited, "Wait should not block when the network is online") +} + +func TestWaitBlocksUntilOnline(t *testing.T) { + s := New() + s.Set(false) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + result := make(chan bool, 1) + go func() { + waited, err := s.Wait(ctx) + if err != nil { + result <- false + return + } + result <- waited + }() + + // Verify Wait is actually blocking while offline + select { + case <-result: + t.Fatal("Wait should block while the network is offline") + case <-time.After(100 * time.Millisecond): + } + + s.Set(true) + + select { + case waited := <-result: + assert.True(t, waited, "Wait should report that it had to wait for the network") + case <-time.After(2 * time.Second): + t.Fatal("Wait should return promptly after the network becomes available") + } +} + +func TestWaitReturnsOnContextCancel(t *testing.T) { + s := New() + s.Set(false) + + ctx, cancel := context.WithCancel(context.Background()) + + result := make(chan error, 1) + go func() { + _, err := s.Wait(ctx) + result <- err + }() + + cancel() + + select { + case err := <-result: + assert.ErrorIs(t, err, context.Canceled) + case <-time.After(2 * time.Second): + t.Fatal("Wait should return promptly after context cancellation") + } +} + +func TestWaitWakesAllWaiters(t *testing.T) { + s := New() + s.Set(false) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + const waiters = 10 + var wg sync.WaitGroup + results := make(chan bool, waiters) + for i := 0; i < waiters; i++ { + wg.Add(1) + go func() { + defer wg.Done() + waited, err := s.Wait(ctx) + if err != nil { + results <- false + return + } + results <- waited + }() + } + + time.Sleep(100 * time.Millisecond) + s.Set(true) + wg.Wait() + + close(results) + count := 0 + for waited := range results { + assert.True(t, waited, "every waiter should report that it waited") + count++ + } + assert.Equal(t, waiters, count, "all waiters should have returned") +} + +func TestNilStateReadsAreNoops(t *testing.T) { + var s *State + + assert.True(t, s.IsOnline(), "nil State should report online") + + waited, err := s.Wait(context.Background()) + require.NoError(t, err) + assert.False(t, waited, "nil State's Wait should not block") +} + +func TestConcurrentSetAndWait(t *testing.T) { + s := New() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + var wg sync.WaitGroup + for i := 0; i < 4; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 100; j++ { + s.Set(j%2 == 0) + s.IsOnline() + } + }() + } + for i := 0; i < 4; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < 100; j++ { + if _, err := s.Wait(ctx); err != nil { + return + } + } + }() + } + + wg.Wait() +} diff --git a/client/netevents/sweep/quick_retry.go b/client/netevents/sweep/quick_retry.go new file mode 100644 index 000000000..1e174b20a --- /dev/null +++ b/client/netevents/sweep/quick_retry.go @@ -0,0 +1,39 @@ +package sweep + +import ( + "time" + + "github.com/cenkalti/backoff/v4" + + "github.com/netbirdio/netbird/client/netevents/netstate" +) + +const quickRetryDelay = 200 * time.Millisecond + +type quickRetryBackoff struct { + backoff.BackOff + sweeper *Sweeper + netState *netstate.State + used bool +} + +func newQuickRetryBackoff(bo backoff.BackOff, sweeper *Sweeper, netState *netstate.State) *quickRetryBackoff { + return &quickRetryBackoff{ + BackOff: bo, + sweeper: sweeper, + netState: netState, + } +} + +func (b *quickRetryBackoff) NextBackOff() time.Duration { + if !b.used && b.sweeper.markedRecently() && b.netState.IsOnline() { + b.used = true + return quickRetryDelay + } + return b.BackOff.NextBackOff() +} + +func (b *quickRetryBackoff) Reset() { + b.used = false + b.BackOff.Reset() +} diff --git a/client/netevents/sweep/quick_retry_test.go b/client/netevents/sweep/quick_retry_test.go new file mode 100644 index 000000000..3dadd951c --- /dev/null +++ b/client/netevents/sweep/quick_retry_test.go @@ -0,0 +1,58 @@ +package sweep + +import ( + "context" + "testing" + "time" + + "github.com/cenkalti/backoff/v4" + "github.com/stretchr/testify/assert" +) + +func TestQuickRetryAfterRecentMark(t *testing.T) { + sweeper := New() + sweeper.MarkNetworkChange() + + slow := backoff.NewConstantBackOff(5 * time.Second) + bo := sweeper.QuickRetryBackoff(context.Background(), slow, nil) + + assert.Equal(t, quickRetryDelay, bo.NextBackOff(), "first retry after a mark must be quick") + assert.Equal(t, 5*time.Second, bo.NextBackOff(), "second retry must fall back to the wrapped backoff") + + bo.Reset() + assert.Equal(t, quickRetryDelay, bo.NextBackOff(), "reset must re-arm the quick retry") +} + +func TestQuickRetryWithoutMarkKeepsSpread(t *testing.T) { + sweeper := New() + + slow := backoff.NewConstantBackOff(5 * time.Second) + bo := sweeper.QuickRetryBackoff(context.Background(), slow, nil) + + assert.Equal(t, 5*time.Second, bo.NextBackOff(), "without a mark the wrapped backoff decides") + + sweeper.mu.Lock() + sweeper.lastMark = time.Now().Add(-recentMarkWindow) + sweeper.mu.Unlock() + assert.Equal(t, 5*time.Second, bo.NextBackOff(), "a stale mark must not trigger the quick retry") +} + +func TestQuickRetryNilSweeperPassthrough(t *testing.T) { + var sweeper *Sweeper + + slow := backoff.NewConstantBackOff(5 * time.Second) + bo := sweeper.QuickRetryBackoff(context.Background(), slow, nil) + + assert.Equal(t, backoff.BackOff(slow), bo, "nil sweeper must return the backoff unchanged") +} + +func TestQuickRetryHonorsContext(t *testing.T) { + sweeper := New() + sweeper.MarkNetworkChange() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + bo := sweeper.QuickRetryBackoff(ctx, backoff.NewConstantBackOff(time.Millisecond), nil) + + assert.Equal(t, backoff.Stop, bo.NextBackOff(), "cancelled context must stop the retry loop") +} diff --git a/client/netevents/sweep/sweep.go b/client/netevents/sweep/sweep.go new file mode 100644 index 000000000..52dce92be --- /dev/null +++ b/client/netevents/sweep/sweep.go @@ -0,0 +1,267 @@ +// Package sweep cuts network-bound activity when the OS switches networks: +// a sweep closes the registered connections and aborts the in-flight dials, so +// their owners redial immediately instead of waiting for the old sockets to +// time out. +// +// A nil *Sweeper disables everything: all methods are nil-safe no-ops. +package sweep + +import ( + "context" + "errors" + "net" + "sync" + "time" + + "github.com/cenkalti/backoff/v4" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/netevents/netstate" +) + +// DefaultSweepDelay absorbs network flapping while the OS settles on a +// default network before the stale registrations are cut. +const DefaultSweepDelay = 500 * time.Millisecond + +const recentMarkWindow = 3 * time.Second + +// Config customizes a Sweeper. The zero value applies the defaults. +type Config struct { + // SweepDelay overrides DefaultSweepDelay when positive. + SweepDelay time.Duration +} + +// ErrSwept reports that a dial finished after a network change swept its +// registration. The connection is already closed; the caller must treat it +// as a failed dial and redial on the new network. +var ErrSwept = errors.New("sweep: connection swept by network change") + +// sweepID identifies one registration in a sweeper. Connections and dials +// draw from the same counter, so an id is unique across both registries. +type sweepID uint64 + +type connEntry struct { + conn net.Conn + gen uint64 +} + +// Dial tracks one dial from start to connection registration. It hands the +// dialed connection to the sweeper atomically, so a sweep can never fall +// between the dial finishing and the connection being registered. +type Dial struct { + sweeper *Sweeper + ctx context.Context + cancel context.CancelFunc + id sweepID + done bool // set by a sweep, WrapConn or Release; guarded by sweeper.mu + gen uint64 +} + +// Ctx returns the dial's context. A sweep cancels it, so a dial started on the +// old network aborts instead of waiting out its handshake timeout. +func (d *Dial) Ctx() context.Context { + return d.ctx +} + +// Release ends the dial's registration and cancels its context. It is +// idempotent and safe after WrapConn, so callers can defer it. +func (d *Dial) Release() { + s := d.sweeper + if s == nil { + return + } + + s.mu.Lock() + d.done = true + delete(s.dials, d.id) + s.mu.Unlock() + + d.cancel() +} + +// sweptConn deregisters itself from the sweeper when closed. +type sweptConn struct { + net.Conn + sweeper *Sweeper + id sweepID +} + +func (c *sweptConn) Close() error { + c.sweeper.deregister(c.id) + return c.Conn.Close() +} + +// Sweeper registers live connections and in-flight dials so the +// network-change sweep can cut everything registered before the change. +type Sweeper struct { + mu sync.Mutex + conns map[sweepID]connEntry + dials map[sweepID]*Dial + nextID sweepID + gen uint64 + timer *time.Timer + sweepDelay time.Duration + lastMark time.Time +} + +// New creates an empty sweeper with the default configuration. +func New() *Sweeper { + return NewWithConfig(Config{}) +} + +// NewWithConfig creates an empty sweeper customized by cfg. +func NewWithConfig(cfg Config) *Sweeper { + delay := cfg.SweepDelay + if delay <= 0 { + delay = DefaultSweepDelay + } + return &Sweeper{ + conns: make(map[sweepID]connEntry), + dials: make(map[sweepID]*Dial), + sweepDelay: delay, + } +} + +// StartDial registers an in-flight dial. Dial with Ctx, hand the result to +// WrapConn, and Release the dial when the attempt is over, typically deferred. +func (s *Sweeper) StartDial(ctx context.Context) *Dial { + if s == nil { + return &Dial{ctx: ctx} + } + + ctx, cancel := context.WithCancel(ctx) + d := &Dial{sweeper: s, ctx: ctx, cancel: cancel} + + s.mu.Lock() + d.id = s.nextID + s.nextID++ + d.gen = s.gen + s.dials[d.id] = d + s.mu.Unlock() + + return d +} + +// WrapConn hands conn over to the sweeper. If a sweep ran since StartDial, +// the connection belongs to the old network: it is closed and ErrSwept is +// returned. Otherwise conn is registered against the next sweep and returned +// wrapped, deregistering itself on Close. Call it once, before Release. +func (d *Dial) WrapConn(conn net.Conn) (net.Conn, error) { + s := d.sweeper + if s == nil { + return conn, nil + } + + s.mu.Lock() + if d.done { + s.mu.Unlock() + if err := conn.Close(); err != nil { + log.Debugf("swept dial close error: %v", err) + } + return nil, ErrSwept + } + d.done = true + delete(s.dials, d.id) + id := s.nextID + s.nextID++ + // The conn inherits the dial's generation: the socket was bound to the + // network that was default when the dial started, not when it finished. + s.conns[id] = connEntry{conn: conn, gen: d.gen} + s.mu.Unlock() + + return &sweptConn{Conn: conn, sweeper: s, id: id}, nil +} + +// MarkNetworkChange records that the OS switched networks: everything +// registered so far becomes stale, and a sweep is (re)scheduled after the +// configured delay to cut whatever is still stale by then. Owners that +// redialed in the meantime hold fresh-generation registrations and survive, +// so no cancellation is needed around the sweep. +func (s *Sweeper) MarkNetworkChange() { + if s == nil { + return + } + + s.mu.Lock() + s.gen++ + cutoff := s.gen + s.lastMark = time.Now() + if s.timer != nil { + s.timer.Stop() + } + s.timer = time.AfterFunc(s.sweepDelay, func() { + n := s.sweep(cutoff) + log.Infof("network change sweep: closed %d stale connections", n) + }) + s.mu.Unlock() +} + +// QuickRetryBackoff wraps bo so that after each Reset the first retry comes +// quickly when the disconnect followed a recent network change and the +// network is online. Any other failure keeps bo's spread, so the clients of +// a restarted server still scatter their reconnects. A nil sweeper returns +// bo unchanged. +func (s *Sweeper) QuickRetryBackoff(ctx context.Context, bo backoff.BackOff, netState *netstate.State) backoff.BackOff { + if s == nil { + return bo + } + return backoff.WithContext(newQuickRetryBackoff(bo, s, netState), ctx) +} + +func (s *Sweeper) markedRecently() bool { + if s == nil { + return false + } + s.mu.Lock() + defer s.mu.Unlock() + return !s.lastMark.IsZero() && time.Since(s.lastMark) < recentMarkWindow +} + +// sweep closes the registered connections and aborts the in-flight dials +// older than cutoff, and returns how many connections it closed. A dial +// whose connection was not yet handed to WrapConn is marked, so the late +// WrapConn closes it instead of registering it. +func (s *Sweeper) sweep(cutoff uint64) int { + if s == nil { + return 0 + } + + s.mu.Lock() + var conns []net.Conn + for id, e := range s.conns { + if e.gen < cutoff { + delete(s.conns, id) + conns = append(conns, e.conn) + } + } + var dials []*Dial + for id, d := range s.dials { + if d.gen < cutoff { + d.done = true + delete(s.dials, id) + dials = append(dials, d) + } + } + s.mu.Unlock() + + if len(dials) > 0 { + log.Debugf("aborting %d in-flight dials", len(dials)) + for _, d := range dials { + d.cancel() + } + } + + for _, conn := range conns { + log.Debugf("sweeping connection %s -> %s", conn.LocalAddr(), conn.RemoteAddr()) + if err := conn.Close(); err != nil { + log.Debugf("swept connection close error: %v", err) + } + } + return len(conns) +} + +func (s *Sweeper) deregister(id sweepID) { + s.mu.Lock() + delete(s.conns, id) + s.mu.Unlock() +} diff --git a/client/netevents/sweep/sweep_test.go b/client/netevents/sweep/sweep_test.go new file mode 100644 index 000000000..c162d4c0f --- /dev/null +++ b/client/netevents/sweep/sweep_test.go @@ -0,0 +1,241 @@ +package sweep + +import ( + "context" + "math" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSweepClosesRegisteredConns(t *testing.T) { + sweeper := New() + + c1 := wrap(t, sweeper, connPair(t)) + c2 := wrap(t, sweeper, connPair(t)) + + assert.Equal(t, 2, sweeper.sweepAll(), "both live connections should be closed") + + // The wrappers must report closed now. + buf := make([]byte, 1) + _, err := c1.Read(buf) + assert.Error(t, err, "first connection should be unusable after the sweep") + _, err = c2.Read(buf) + assert.Error(t, err, "second connection should be unusable after the sweep") + + assert.Equal(t, 0, sweeper.sweepAll(), "second sweep should find nothing") +} + +func TestCloseDeregisters(t *testing.T) { + sweeper := New() + + conn := wrap(t, sweeper, connPair(t)) + require.NoError(t, conn.Close()) + + assert.Equal(t, 0, sweeper.sweepAll(), "closed connection must leave the registry") +} + +func TestCloseIsIdempotent(t *testing.T) { + sweeper := New() + + conn := wrap(t, sweeper, connPair(t)) + require.NoError(t, conn.Close()) + assert.Error(t, conn.Close(), "double close surfaces the underlying error but must not panic") +} + +func TestSweepOnlyAffectsOlderConns(t *testing.T) { + sweeper := New() + + _ = wrap(t, sweeper, connPair(t)) + assert.Equal(t, 1, sweeper.sweepAll()) + + // A connection dialed after the sweep must survive until the next one. + _ = wrap(t, sweeper, connPair(t)) + assert.Equal(t, 1, sweeper.sweepAll(), "post-sweep connection belongs to the next sweep") +} + +func TestSweepAbortsInFlightDials(t *testing.T) { + sweeper := New() + + dial := sweeper.StartDial(context.Background()) + defer dial.Release() + + sweeper.sweepAll() + + assert.ErrorIs(t, dial.Ctx().Err(), context.Canceled, "sweep must cancel the in-flight dial context") +} + +func TestReleasedDialIsNotAborted(t *testing.T) { + sweeper := New() + + // Simulate a dial that finished before the sweep. + released := sweeper.StartDial(context.Background()) + released.Release() + + // A dial still in flight during the sweep. + pending := sweeper.StartDial(context.Background()) + defer pending.Release() + + sweeper.sweepAll() + assert.ErrorIs(t, pending.Ctx().Err(), context.Canceled, "pending dial must be aborted") +} + +func TestSweepBetweenDialAndHandoffClosesConn(t *testing.T) { + sweeper := New() + + dial := sweeper.StartDial(context.Background()) + defer dial.Release() + + // The dial succeeds on the old network, then the sweep lands before the + // connection is handed over. + conn := connPair(t) + assert.Equal(t, 0, sweeper.sweepAll(), "the connection is not registered yet") + + wrapped, err := dial.WrapConn(conn) + require.ErrorIs(t, err, ErrSwept) + require.Nil(t, wrapped) + + buf := make([]byte, 1) + _, err = conn.Read(buf) + assert.Error(t, err, "the old-network connection must be closed, not leaked") + + assert.Equal(t, 0, sweeper.sweepAll(), "nothing may leak into the next sweep") +} + +func TestMarkNetworkChangeSparesFreshConns(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 10 * time.Millisecond}) + + stale := wrap(t, sweeper, connPair(t)) + sweeper.MarkNetworkChange() + _ = wrap(t, sweeper, connPair(t)) + + _ = stale.SetReadDeadline(time.Now().Add(time.Second)) + buf := make([]byte, 1) + _, err := stale.Read(buf) + require.ErrorIs(t, err, net.ErrClosed, "stale connection must be closed by the delayed sweep") + + assert.Equal(t, 1, sweeper.sweepAll(), "the fresh connection must survive the stale sweep") +} + +func TestMarkNetworkChangeAbortsStaleDials(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 10 * time.Millisecond}) + + stale := sweeper.StartDial(context.Background()) + defer stale.Release() + sweeper.MarkNetworkChange() + fresh := sweeper.StartDial(context.Background()) + defer fresh.Release() + + assert.Eventually(t, func() bool { + return stale.Ctx().Err() != nil + }, time.Second, 5*time.Millisecond, "stale dial must be aborted by the delayed sweep") + assert.NoError(t, fresh.Ctx().Err(), "post-mark dial must not be aborted") +} + +func TestConnInheritsDialGeneration(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 20 * time.Millisecond}) + + // The dial starts before the network change but completes after it: the + // socket is bound to the old network, so the sweep must still cut it. + dial := sweeper.StartDial(context.Background()) + defer dial.Release() + sweeper.MarkNetworkChange() + + wrapped, err := dial.WrapConn(connPair(t)) + require.NoError(t, err) + + _ = wrapped.SetReadDeadline(time.Now().Add(time.Second)) + buf := make([]byte, 1) + _, err = wrapped.Read(buf) + require.ErrorIs(t, err, net.ErrClosed, "old-generation connection must be swept") +} + +func TestRepeatedMarksCoalesce(t *testing.T) { + sweeper := NewWithConfig(Config{SweepDelay: 20 * time.Millisecond}) + + first := wrap(t, sweeper, connPair(t)) + sweeper.MarkNetworkChange() + second := wrap(t, sweeper, connPair(t)) + sweeper.MarkNetworkChange() + _ = wrap(t, sweeper, connPair(t)) + + buf := make([]byte, 1) + for _, conn := range []net.Conn{first, second} { + _ = conn.SetReadDeadline(time.Now().Add(time.Second)) + _, err := conn.Read(buf) + require.ErrorIs(t, err, net.ErrClosed, "every pre-mark connection must be swept by the rescheduled sweep") + } + assert.Equal(t, 1, sweeper.sweepAll(), "only the newest-generation connection may remain") +} + +func TestNilSweeperIsNoop(t *testing.T) { + var sweeper *Sweeper + + conn := connPair(t) + dial := sweeper.StartDial(context.Background()) + defer dial.Release() + + wrapped, err := dial.WrapConn(conn) + require.NoError(t, err) + assert.Equal(t, conn, wrapped, "nil sweeper must return the conn unchanged") + assert.NoError(t, dial.Ctx().Err(), "nil sweeper must not cancel the dial context") + assert.Equal(t, 0, sweeper.sweepAll(), "nil sweeper closes nothing") +} + +// wrap registers conn with the sweeper through a completed dial. +func wrap(t *testing.T, sweeper *Sweeper, conn net.Conn) net.Conn { + t.Helper() + + dial := sweeper.StartDial(context.Background()) + defer dial.Release() + + wrapped, err := dial.WrapConn(conn) + require.NoError(t, err) + return wrapped +} + +// connPair dials a loopback TCP connection and keeps the accepted peer open +// until the test ends: a peer that closed early would make the connection +// unreadable on its own, so a read error after the sweep would prove nothing. +func connPair(t *testing.T) net.Conn { + t.Helper() + + l, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { + if err := l.Close(); err != nil { + t.Logf("listener close error: %v", err) + } + }) + + accepted := make(chan net.Conn, 1) + go func() { + conn, err := l.Accept() + if err != nil { + close(accepted) + return + } + accepted <- conn + }() + + conn, err := net.Dial("tcp", l.Addr().String()) + require.NoError(t, err) + + peer, ok := <-accepted + require.True(t, ok, "listener must accept the dialed connection") + t.Cleanup(func() { + if err := peer.Close(); err != nil { + t.Logf("peer close error: %v", err) + } + }) + + return conn +} + +// sweepAll cuts every registration regardless of generation. +func (s *Sweeper) sweepAll() int { + return s.sweep(math.MaxUint64) +} diff --git a/client/proto/daemon.pb.go b/client/proto/daemon.pb.go index 6fbb09958..b438a310a 100644 --- a/client/proto/daemon.pb.go +++ b/client/proto/daemon.pb.go @@ -2771,14 +2771,23 @@ func (x *ForwardingRulesResponse) GetRules() []*ForwardingRule { // DebugBundler type DebugBundleRequest struct { - state protoimpl.MessageState `protogen:"open.v1"` - Anonymize bool `protobuf:"varint,1,opt,name=anonymize,proto3" json:"anonymize,omitempty"` - SystemInfo bool `protobuf:"varint,3,opt,name=systemInfo,proto3" json:"systemInfo,omitempty"` - UploadURL string `protobuf:"bytes,4,opt,name=uploadURL,proto3" json:"uploadURL,omitempty"` - LogFileCount uint32 `protobuf:"varint,5,opt,name=logFileCount,proto3" json:"logFileCount,omitempty"` - CliVersion string `protobuf:"bytes,6,opt,name=cliVersion,proto3" json:"cliVersion,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache + state protoimpl.MessageState `protogen:"open.v1"` + Anonymize bool `protobuf:"varint,1,opt,name=anonymize,proto3" json:"anonymize,omitempty"` + SystemInfo bool `protobuf:"varint,3,opt,name=systemInfo,proto3" json:"systemInfo,omitempty"` + UploadURL string `protobuf:"bytes,4,opt,name=uploadURL,proto3" json:"uploadURL,omitempty"` + LogFileCount uint32 `protobuf:"varint,5,opt,name=logFileCount,proto3" json:"logFileCount,omitempty"` + CliVersion string `protobuf:"bytes,6,opt,name=cliVersion,proto3" json:"cliVersion,omitempty"` + // uploadInsecure allows uploading to an http endpoint or one with an + // untrusted TLS certificate. Restricted to privileged callers; for + // self-hosted upload servers. + UploadInsecure bool `protobuf:"varint,7,opt,name=uploadInsecure,proto3" json:"uploadInsecure,omitempty"` + // anonymizeLevel selects how much the anonymizer redacts: "default" + // (or empty) keeps internal IP ranges, "strict" also anonymizes them. + // Unknown values are treated as "strict". Only meaningful with anonymize; + // "strict" implies it. + AnonymizeLevel string `protobuf:"bytes,8,opt,name=anonymizeLevel,proto3" json:"anonymizeLevel,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache } func (x *DebugBundleRequest) Reset() { @@ -2846,6 +2855,20 @@ func (x *DebugBundleRequest) GetCliVersion() string { return "" } +func (x *DebugBundleRequest) GetUploadInsecure() bool { + if x != nil { + return x.UploadInsecure + } + return false +} + +func (x *DebugBundleRequest) GetAnonymizeLevel() string { + if x != nil { + return x.AnonymizeLevel + } + return "" +} + type DebugBundleResponse struct { state protoimpl.MessageState `protogen:"open.v1"` Path string `protobuf:"bytes,1,opt,name=path,proto3" json:"path,omitempty"` @@ -5605,9 +5628,13 @@ func (x *GetPeerSSHHostKeyResponse) GetFound() bool { type RequestJWTAuthRequest struct { state protoimpl.MessageState `protogen:"open.v1"` // hint for OIDC login_hint parameter (typically email address) - Hint *string `protobuf:"bytes,1,opt,name=hint,proto3,oneof" json:"hint,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache + Hint *string `protobuf:"bytes,1,opt,name=hint,proto3,oneof" json:"hint,omitempty"` + // hasGraphicalSession tells the daemon that the caller has a graphical session, + // which decides whether PKCE or the device code flow is preferred. The daemon + // cannot detect this itself: it does not inherit the session environment. + HasGraphicalSession bool `protobuf:"varint,2,opt,name=hasGraphicalSession,proto3" json:"hasGraphicalSession,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache } func (x *RequestJWTAuthRequest) Reset() { @@ -5647,6 +5674,13 @@ func (x *RequestJWTAuthRequest) GetHint() string { return "" } +func (x *RequestJWTAuthRequest) GetHasGraphicalSession() bool { + if x != nil { + return x.HasGraphicalSession + } + return false +} + // RequestJWTAuthResponse contains authentication flow information type RequestJWTAuthResponse struct { state protoimpl.MessageState `protogen:"open.v1"` @@ -5871,9 +5905,13 @@ type RequestExtendAuthSessionRequest struct { state protoimpl.MessageState `protogen:"open.v1"` // Optional OIDC login_hint (typically the user's email) to pre-fill the // IdP login form. - Hint *string `protobuf:"bytes,1,opt,name=hint,proto3,oneof" json:"hint,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache + Hint *string `protobuf:"bytes,1,opt,name=hint,proto3,oneof" json:"hint,omitempty"` + // hasGraphicalSession tells the daemon that the caller has a graphical session, + // which decides whether PKCE or the device code flow is preferred. The daemon + // cannot detect this itself: it does not inherit the session environment. + HasGraphicalSession bool `protobuf:"varint,2,opt,name=hasGraphicalSession,proto3" json:"hasGraphicalSession,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache } func (x *RequestExtendAuthSessionRequest) Reset() { @@ -5913,6 +5951,13 @@ func (x *RequestExtendAuthSessionRequest) GetHint() string { return "" } +func (x *RequestExtendAuthSessionRequest) GetHasGraphicalSession() bool { + if x != nil { + return x.HasGraphicalSession + } + return false +} + // RequestExtendAuthSessionResponse carries the verification URI the UI // should open in a browser. The daemon retains the flow state and resolves // it via WaitExtendAuthSession. @@ -7242,7 +7287,7 @@ const file_daemon_proto_rawDesc = "" + "\x12translatedHostname\x18\x04 \x01(\tR\x12translatedHostname\x128\n" + "\x0etranslatedPort\x18\x05 \x01(\v2\x10.daemon.PortInfoR\x0etranslatedPort\"G\n" + "\x17ForwardingRulesResponse\x12,\n" + - "\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\xb4\x01\n" + + "\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\x84\x02\n" + "\x12DebugBundleRequest\x12\x1c\n" + "\tanonymize\x18\x01 \x01(\bR\tanonymize\x12\x1e\n" + "\n" + @@ -7252,7 +7297,9 @@ const file_daemon_proto_rawDesc = "" + "\flogFileCount\x18\x05 \x01(\rR\flogFileCount\x12\x1e\n" + "\n" + "cliVersion\x18\x06 \x01(\tR\n" + - "cliVersion\"}\n" + + "cliVersion\x12&\n" + + "\x0euploadInsecure\x18\a \x01(\bR\x0euploadInsecure\x12&\n" + + "\x0eanonymizeLevel\x18\b \x01(\tR\x0eanonymizeLevel\"}\n" + "\x13DebugBundleResponse\x12\x12\n" + "\x04path\x18\x01 \x01(\tR\x04path\x12 \n" + "\vuploadedKey\x18\x02 \x01(\tR\vuploadedKey\x120\n" + @@ -7478,9 +7525,10 @@ const file_daemon_proto_rawDesc = "" + "sshHostKey\x12\x16\n" + "\x06peerIP\x18\x02 \x01(\tR\x06peerIP\x12\x1a\n" + "\bpeerFQDN\x18\x03 \x01(\tR\bpeerFQDN\x12\x14\n" + - "\x05found\x18\x04 \x01(\bR\x05found\"9\n" + + "\x05found\x18\x04 \x01(\bR\x05found\"k\n" + "\x15RequestJWTAuthRequest\x12\x17\n" + - "\x04hint\x18\x01 \x01(\tH\x00R\x04hint\x88\x01\x01B\a\n" + + "\x04hint\x18\x01 \x01(\tH\x00R\x04hint\x88\x01\x01\x120\n" + + "\x13hasGraphicalSession\x18\x02 \x01(\bR\x13hasGraphicalSessionB\a\n" + "\x05_hint\"\x9a\x02\n" + "\x16RequestJWTAuthResponse\x12(\n" + "\x0fverificationURI\x18\x01 \x01(\tR\x0fverificationURI\x128\n" + @@ -7500,9 +7548,10 @@ const file_daemon_proto_rawDesc = "" + "\x14WaitJWTTokenResponse\x12\x14\n" + "\x05token\x18\x01 \x01(\tR\x05token\x12\x1c\n" + "\ttokenType\x18\x02 \x01(\tR\ttokenType\x12\x1c\n" + - "\texpiresIn\x18\x03 \x01(\x03R\texpiresIn\"C\n" + + "\texpiresIn\x18\x03 \x01(\x03R\texpiresIn\"u\n" + "\x1fRequestExtendAuthSessionRequest\x12\x17\n" + - "\x04hint\x18\x01 \x01(\tH\x00R\x04hint\x88\x01\x01B\a\n" + + "\x04hint\x18\x01 \x01(\tH\x00R\x04hint\x88\x01\x01\x120\n" + + "\x13hasGraphicalSession\x18\x02 \x01(\bR\x13hasGraphicalSessionB\a\n" + "\x05_hint\"\xe0\x01\n" + " RequestExtendAuthSessionResponse\x12(\n" + "\x0fverificationURI\x18\x01 \x01(\tR\x0fverificationURI\x128\n" + diff --git a/client/proto/daemon.proto b/client/proto/daemon.proto index 8d5294eb7..a3e3f4500 100644 --- a/client/proto/daemon.proto +++ b/client/proto/daemon.proto @@ -536,6 +536,15 @@ message DebugBundleRequest { string uploadURL = 4; uint32 logFileCount = 5; string cliVersion = 6; + // uploadInsecure allows uploading to an http endpoint or one with an + // untrusted TLS certificate. Restricted to privileged callers; for + // self-hosted upload servers. + bool uploadInsecure = 7; + // anonymizeLevel selects how much the anonymizer redacts: "default" + // (or empty) keeps internal IP ranges, "strict" also anonymizes them. + // Unknown values are treated as "strict". Only meaningful with anonymize; + // "strict" implies it. + string anonymizeLevel = 8; } message DebugBundleResponse { @@ -885,6 +894,10 @@ message GetPeerSSHHostKeyResponse { message RequestJWTAuthRequest { // hint for OIDC login_hint parameter (typically email address) optional string hint = 1; + // hasGraphicalSession tells the daemon that the caller has a graphical session, + // which decides whether PKCE or the device code flow is preferred. The daemon + // cannot detect this itself: it does not inherit the session environment. + bool hasGraphicalSession = 2; } // RequestJWTAuthResponse contains authentication flow information @@ -928,6 +941,10 @@ message RequestExtendAuthSessionRequest { // Optional OIDC login_hint (typically the user's email) to pre-fill the // IdP login form. optional string hint = 1; + // hasGraphicalSession tells the daemon that the caller has a graphical session, + // which decides whether PKCE or the device code flow is preferred. The daemon + // cannot detect this itself: it does not inherit the session environment. + bool hasGraphicalSession = 2; } // RequestExtendAuthSessionResponse carries the verification URI the UI diff --git a/client/server/debug.go b/client/server/debug.go index 0b6ac4b53..8f4a506b4 100644 --- a/client/server/debug.go +++ b/client/server/debug.go @@ -7,18 +7,63 @@ import ( "context" "errors" "fmt" + "path/filepath" "runtime/pprof" + "strings" + "time" log "github.com/sirupsen/logrus" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + "github.com/netbirdio/netbird/client/anonymize" "github.com/netbirdio/netbird/client/internal/debug" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" mgmProto "github.com/netbirdio/netbird/shared/management/proto" "github.com/netbirdio/netbird/version" ) // DebugBundle creates a debug bundle and returns the location. -func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) (resp *proto.DebugBundleResponse, err error) { +func (s *Server) DebugBundle(callerCtx context.Context, req *proto.DebugBundleRequest) (resp *proto.DebugBundleResponse, err error) { + if err := requirePrivilegeForUploadURL(callerCtx, req.GetUploadURL(), req.GetUploadInsecure()); err != nil { + return nil, err + } + + // The UI log is opened as whoever asked for this bundle, so a caller only + // collects a log it owns (privileged callers excepted). ok is false on a + // socket that carries no identity, which skips the UI log. + callerID, callerIdentified := ipcauth.CallerIdentity(callerCtx) + + path, managementURL, err := s.generateDebugBundle(req, uiLogOpener(callerID, callerIdentified)) + if err != nil { + return nil, err + } + + if req.GetUploadURL() == "" { + return &proto.DebugBundleResponse{Path: path}, nil + } + + // The upload runs without s.mutex held: it does network I/O to a possibly + // slow destination and must not block the other RPCs that take the lock. The + // bounded context is a backstop against a hung connection. + uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() + key, err := debug.UploadDebugBundle(uploadCtx, req.GetUploadURL(), managementURL, path, req.GetUploadInsecure()) + if err != nil { + log.Errorf("failed to upload debug bundle to %s: %v", req.GetUploadURL(), err) + return &proto.DebugBundleResponse{Path: path, UploadFailureReason: err.Error()}, nil + } + + log.Infof("debug bundle uploaded to %s with key %s", req.GetUploadURL(), key) + + return &proto.DebugBundleResponse{Path: path, UploadedKey: key}, nil +} + +// generateDebugBundle builds the bundle under s.mutex and returns its path plus +// the management URL captured under the lock, so the caller can run the upload +// without holding the lock. +func (s *Server) generateDebugBundle(req *proto.DebugBundleRequest, uiOpener debug.LogOpener) (path string, managementURL string, err error) { s.mutex.Lock() defer s.mutex.Unlock() @@ -68,6 +113,7 @@ func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) ( SyncResponse: syncResponse, LogPath: s.logFile, UILogPath: s.uiLogPath, + UILogOpener: uiOpener, CPUProfile: cpuProfileData, CapturePath: capturePath, RefreshStatus: refreshStatus, @@ -77,28 +123,22 @@ func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) ( }, debug.BundleConfig{ Anonymize: req.GetAnonymize(), + AnonymizeLevel: anonymize.ParseLevel(req.GetAnonymizeLevel()), IncludeSystemInfo: req.GetSystemInfo(), LogFileCount: req.GetLogFileCount(), }, ) - path, err := bundleGenerator.Generate() + path, err = bundleGenerator.Generate() if err != nil { - return nil, fmt.Errorf("generate debug bundle: %w", err) + return "", "", fmt.Errorf("generate debug bundle: %w", err) } - if req.GetUploadURL() == "" { - return &proto.DebugBundleResponse{Path: path}, nil - } - key, err := debug.UploadDebugBundle(context.Background(), req.GetUploadURL(), s.config.ManagementURL.String(), path) - if err != nil { - log.Errorf("failed to upload debug bundle to %s: %v", req.GetUploadURL(), err) - return &proto.DebugBundleResponse{Path: path, UploadFailureReason: err.Error()}, nil + if s.config != nil && s.config.ManagementURL != nil { + managementURL = s.config.ManagementURL.String() } - log.Infof("debug bundle uploaded to %s with key %s", req.GetUploadURL(), key) - - return &proto.DebugBundleResponse{Path: path, UploadedKey: key}, nil + return path, managementURL, nil } // GetLogLevel gets the current logging level for the server. @@ -138,12 +178,34 @@ func (s *Server) SetLogLevel(_ context.Context, req *proto.SetLogLevelRequest) ( // RegisterUILog records the desktop UI's absolute log path so DebugBundle can // collect the GUI log. The daemon runs as root and can't resolve the user's // config dir, so the UI reports it. Last-writer-wins (one UI per socket). -func (s *Server) RegisterUILog(_ context.Context, req *proto.RegisterUILogRequest) (*proto.RegisterUILogResponse, error) { +// +// The path arrives over an IPC any local user can reach and is later opened by +// a root daemon, so it is constrained to the file name the UI writes and to a +// local absolute path. Authorization happens when DebugBundle opens it: the +// bundle refuses a file its requester does not own. A caller the daemon cannot +// identify cannot register a path at all. +func (s *Server) RegisterUILog(callerCtx context.Context, req *proto.RegisterUILogRequest) (*proto.RegisterUILogResponse, error) { + if _, ok := ipcauth.CallerIdentity(callerCtx); !ok { + return nil, gstatus.Error(codes.PermissionDenied, + "registering a UI log path requires a control channel that carries the caller's identity") + } + + path := filepath.Clean(req.GetPath()) + if !filepath.IsAbs(path) || filepath.Base(path) != uiLogFileName { + return nil, gstatus.Errorf(codes.InvalidArgument, "UI log path must be an absolute path ending in %s", uiLogFileName) + } + // filepath.IsAbs accepts a Windows UNC path (\\host\share\...) and a device + // path (\\.\, \\?\); opening one would make the root daemon reach a remote + // or device namespace. Require a plain local path. + if strings.HasPrefix(path, `\\`) { + return nil, gstatus.Error(codes.InvalidArgument, "UI log path must be a local path, not a UNC or device path") + } + s.mutex.Lock() defer s.mutex.Unlock() - s.uiLogPath = req.GetPath() - log.Infof("registered UI log path: %s", s.uiLogPath) + s.uiLogPath = path + log.Infof("registered UI log path %s", s.uiLogPath) return &proto.RegisterUILogResponse{}, nil } diff --git a/client/server/debug_gate.go b/client/server/debug_gate.go new file mode 100644 index 000000000..983a13aaf --- /dev/null +++ b/client/server/debug_gate.go @@ -0,0 +1,99 @@ +//go:build !android && !ios + +package server + +import ( + "context" + "fmt" + "net/url" + "os" + "strings" + + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/configs" + "github.com/netbirdio/netbird/client/internal/debug" + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/upload-server/types" +) + +// uiLogFileName is the only file name the daemon accepts as a UI log path. The +// UI (writer), this validation, and the bundle collector all read it from +// configs so they cannot drift. +const uiLogFileName = configs.UILogFile + +// uiLogOpener opens the registered UI log, and its rotated siblings, on behalf +// of the caller requesting the bundle: OpenOwnedFile then collects the log only +// when that caller owns it (or is privileged). identified is false on a socket +// that carries no caller identity, in which case nothing is opened. +func uiLogOpener(id ipcauth.Identity, identified bool) debug.LogOpener { + return func(path string) (*os.File, error) { + if !identified { + return nil, fmt.Errorf("bundle requester has no verified identity") + } + return ipcauth.OpenOwnedFile(id, path) + } +} + +// requirePrivilegeForUploadURL restricts where the daemon may send a debug +// bundle. The bundle holds the daemon's own logs and state, and the daemon +// fetches the upload URL itself, so an unrestricted endpoint turns the daemon +// into both an exfiltration channel and a request forwarder that reaches +// services only it can talk to. +// +// The upload service NetBird publishes is open to any caller, since that is what +// the CLI and the desktop UI use. Any other endpoint, self-hosted upload servers +// included, requires a privileged caller. Plaintext is refused for everyone: the +// daemon fetches the URL and then PUTs the bundle to whatever that fetch returns, +// so an http hop is a place to intercept the bundle or the redirect. +// +// insecure relaxes transport security (http, or an untrusted TLS certificate) +// for a self-hosted server. It weakens a root-privileged upload, so it is +// refused for an unprivileged caller regardless of the host. +func requirePrivilegeForUploadURL(ctx context.Context, rawURL string, insecure bool) error { + if rawURL == "" { + return nil + } + + parsed, err := url.Parse(rawURL) + if err != nil { + return gstatus.Errorf(codes.InvalidArgument, "parse upload URL: %v", err) + } + + // --insecure relaxes https to http or an untrusted certificate; it does not + // widen the URL to arbitrary schemes, so a host and http/https are required + // before the insecure branch takes over. + if parsed.Host == "" || (parsed.Scheme != "https" && parsed.Scheme != "http") { + return gstatus.Errorf(codes.InvalidArgument, "upload URL must be http or https with a host") + } + + if insecure { + return denyPrivileged(ctx, + "uploading a debug bundle without transport security (--upload-bundle-insecure)", + ipcauth.ElevatedCommand("netbird debug bundle -U --upload-bundle-insecure --upload-bundle-url ")) + } + + if parsed.Scheme != "https" { + return gstatus.Errorf(codes.InvalidArgument, "upload URL must use https, got scheme %q", parsed.Scheme) + } + + if isDefaultUploadService(parsed) { + return nil + } + + return denyPrivileged(ctx, + "uploading a debug bundle to an upload service other than the default one", + ipcauth.ElevatedCommand("netbird debug bundle -U --upload-bundle-url ")) +} + +// isDefaultUploadService reports whether the URL points at the upload service +// NetBird runs. Only the host is compared: the service's path may differ between +// releases, and the host is what decides who receives the bundle. +func isDefaultUploadService(parsed *url.URL) bool { + defaultURL, err := url.Parse(types.DefaultBundleURL) + if err != nil { + return false + } + return parsed.Scheme == defaultURL.Scheme && strings.EqualFold(parsed.Host, defaultURL.Host) +} diff --git a/client/server/debug_gate_test.go b/client/server/debug_gate_test.go new file mode 100644 index 000000000..e958fc581 --- /dev/null +++ b/client/server/debug_gate_test.go @@ -0,0 +1,157 @@ +//go:build !android && !ios + +package server + +import ( + "os" + "path/filepath" + "runtime" + "testing" + + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/proto" + "github.com/netbirdio/netbird/upload-server/types" +) + +func TestRegisterUILogRefusesUnidentifiedCaller(t *testing.T) { + s := &Server{} + + _, err := s.RegisterUILog(noIdentityCtx(), &proto.RegisterUILogRequest{ + Path: filepath.Join(t.TempDir(), uiLogFileName), + }) + + if gstatus.Code(err) != codes.PermissionDenied { + t.Fatalf("code = %v, want PermissionDenied", gstatus.Code(err)) + } +} + +func TestRegisterUILogRefusesForeignPath(t *testing.T) { + secret := "/etc/shadow" + if runtime.GOOS == "windows" { + secret = `C:\Windows\System32\config\SAM` + } + + tests := []struct { + name string + path string + }{ + {"empty", ""}, + {"relative", filepath.Join("netbird", uiLogFileName)}, + {"another file", secret}, + {"directory of the log", t.TempDir()}, + {"unc path", `\\attacker\share\` + uiLogFileName}, + {"device path", `\\.\C:\` + uiLogFileName}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + s := &Server{} + + _, err := s.RegisterUILog(userCtx(), &proto.RegisterUILogRequest{Path: tc.path}) + + if gstatus.Code(err) != codes.InvalidArgument { + t.Fatalf("code = %v, want InvalidArgument", gstatus.Code(err)) + } + if s.uiLogPath != "" { + t.Fatalf("path %q was recorded despite the refusal", s.uiLogPath) + } + }) + } +} + +func TestRegisterUILogRecordsPath(t *testing.T) { + s := &Server{} + path := filepath.Join(t.TempDir(), uiLogFileName) + + if _, err := s.RegisterUILog(userCtx(), &proto.RegisterUILogRequest{Path: path}); err != nil { + t.Fatalf("register: %v", err) + } + + if s.uiLogPath != path { + t.Fatalf("path = %q, want %q", s.uiLogPath, path) + } +} + +// The UI log is opened as the bundle requester, so a second local user cannot +// collect a log they do not own, and an unidentified requester collects nothing. +func TestUILogOpenerBindsToRequester(t *testing.T) { + path := filepath.Join(t.TempDir(), uiLogFileName) + if err := os.WriteFile(path, []byte("log line"), 0600); err != nil { + t.Fatalf("write log: %v", err) + } + + // A different unprivileged user than the file's owner: refused. + if _, err := uiLogOpener(unprivilegedIdentity(), true)(path); err == nil { + t.Fatal("expected a file the requester does not own to be refused") + } + + // No verified identity: refused. + if _, err := uiLogOpener(ipcauth.Identity{}, false)(path); err == nil { + t.Fatal("expected an unidentified requester to be refused") + } + + // The requester that owns the file: allowed. The test process created it, so + // its own identity is the owner (and a privileged runner is exempt anyway). + owner, err := ipcauth.CurrentProcessIdentity() + if err != nil { + t.Fatalf("current identity: %v", err) + } + f, err := uiLogOpener(owner, true)(path) + if err != nil { + t.Fatalf("expected the owning requester to be allowed, got %v", err) + } + _ = f.Close() +} + +func TestRequirePrivilegeForUploadURL(t *testing.T) { + tests := []struct { + name string + url string + insecure bool + unprivOK bool + invalid bool + rootAlso bool + }{ + {name: "no upload", url: "", unprivOK: true}, + {name: "default service", url: types.DefaultBundleURL, unprivOK: true}, + {name: "default service, other path", url: "https://upload.debug.netbird.io/other", unprivOK: true}, + {name: "loopback exfiltration endpoint", url: "https://127.0.0.1:8080/upload-url", rootAlso: true}, + {name: "custom upload service", url: "https://attacker.example/upload-url", rootAlso: true}, + {name: "plaintext default host", url: "http://upload.debug.netbird.io/upload-url", invalid: true}, + {name: "plaintext custom host", url: "http://attacker.example/upload-url", invalid: true}, + {name: "unsupported scheme", url: "file:///etc/shadow", invalid: true}, + // insecure relaxes transport security; privileged only, whatever the host. + {name: "insecure http custom", url: "http://selfhosted.local/upload-url", insecure: true, rootAlso: true}, + {name: "insecure https custom", url: "https://selfhosted.local/upload-url", insecure: true, rootAlso: true}, + {name: "insecure default host", url: types.DefaultBundleURL, insecure: true, rootAlso: true}, + // --insecure must not widen the URL to non-http(s) schemes or a hostless URL. + {name: "insecure file scheme", url: "file:///etc/shadow", insecure: true, invalid: true}, + {name: "insecure hostless", url: "https:///upload-url", insecure: true, invalid: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + err := requirePrivilegeForUploadURL(userCtx(), tc.url, tc.insecure) + + switch { + case tc.invalid: + if gstatus.Code(err) != codes.InvalidArgument { + t.Fatalf("code = %v, want InvalidArgument", gstatus.Code(err)) + } + return + case tc.unprivOK: + assertAllowed(t, err) + return + default: + assertDenied(t, err) + } + + if tc.rootAlso { + assertAllowed(t, requirePrivilegeForUploadURL(rootCtx(), tc.url, tc.insecure)) + } + }) + } +} diff --git a/client/server/lock_order_test.go b/client/server/lock_order_test.go new file mode 100644 index 000000000..457e3db34 --- /dev/null +++ b/client/server/lock_order_test.go @@ -0,0 +1,51 @@ +package server + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/proto" +) + +// The daemon takes guardedConfigMu before s.mutex. authorizeAndPrepareLogin +// takes s.mutex while holding guardedConfigMu, so a SetConfig that grabbed +// s.mutex first and then waited for guardedConfigMu would deadlock the daemon +// against a concurrent login: two unprivileged IPC calls are enough. +// +// The held guardedConfigMu below stands in for that login. While SetConfig waits +// for it, s.mutex must stay free, otherwise the login waiting for s.mutex could +// never release guardedConfigMu. +func TestSetConfig_TakesGuardedConfigMuBeforeServerMutex(t *testing.T) { + s, ctx, profName, username, _ := setupServerWithProfile(t) + + s.guardedConfigMu.Lock() + + done := make(chan error, 1) + go func() { + _, err := s.SetConfig(ctx, &proto.SetConfigRequest{ + ProfileName: profName, + Username: username, + }) + done <- err + }() + + require.Never(t, func() bool { + if !s.mutex.TryLock() { + return true + } + s.mutex.Unlock() + return false + }, 500*time.Millisecond, 10*time.Millisecond, + "SetConfig held s.mutex while waiting for guardedConfigMu, which deadlocks against a concurrent login") + + s.guardedConfigMu.Unlock() + + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("SetConfig did not finish after guardedConfigMu was released") + } +} diff --git a/client/server/login_gate_test.go b/client/server/login_gate_test.go new file mode 100644 index 000000000..de62a8180 --- /dev/null +++ b/client/server/login_gate_test.go @@ -0,0 +1,127 @@ +package server + +import ( + "context" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/proto" +) + +// A refused login must not leave the profile switched. Login can both switch +// profiles and carry the guarded config fields, so the gate has to run before the +// switch: otherwise a caller whose change is refused still gets the side effect of +// activating whichever profile the request named. +func TestLogin_RefusedChangeLeavesTheProfileAlone(t *testing.T) { + s, _, activeProfile, username, _ := setupServerWithProfile(t) + + // Login reads process state off the daemon's root context. + s.rootCtx = internal.CtxInitState(context.Background()) + + // A second profile that runs the SSH server, which is what makes repointing + // its management binding a privileged change. + target := "ssh-enabled" + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: filepath.Join(profilemanager.DefaultConfigPathDir, target+".json"), + ManagementURL: "https://api.netbird.io:443", + ServerSSHAllowed: boolPtr(true), + }) + require.NoError(t, err) + + _, err = s.Login(userCtx(), &proto.LoginRequest{ + ProfileName: &target, + Username: &username, + ManagementUrl: "https://mgmt.attacker.example:443", + }) + require.Error(t, err, "an unprivileged caller must not move the management URL of an SSH-enabled profile") + require.Equal(t, codes.PermissionDenied, gstatus.Code(err), "want a privilege refusal, got %v", err) + + active, err := s.profileManager.GetActiveProfileState() + require.NoError(t, err) + require.Equal(t, profilemanager.ID(activeProfile), active.ID, + "the refused login switched the active profile anyway") +} + +// A caller whose change becomes privileged only after its first check must be +// refused without having cancelled a login or switched profiles: the first check is +// unsynchronized, so the SSH server can be enabled by a concurrent privileged +// request in between, and the authoritative check happens before any side effect. +func TestLogin_ChangeThatBecomesPrivilegedMidRequestHasNoSideEffects(t *testing.T) { + s, _, activeProfile, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + // The target profile has SSH off, so the first check lets the request through. + target := "ssh-later" + targetPath := filepath.Join(profilemanager.DefaultConfigPathDir, target+".json") + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: targetPath, + ManagementURL: "https://api.netbird.io:443", + ServerSSHAllowed: boolPtr(false), + }) + require.NoError(t, err) + + cancelled := false + s.actCancel = func() { cancelled = true } + + // Stand in for a privileged SetConfig that enables the SSH server between the + // two checks, which is the interleaving the lock has to make safe. + afterLoginPreCheck = func() { + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: targetPath, + ServerSSHAllowed: boolPtr(true), + }) + require.NoError(t, err) + } + t.Cleanup(func() { afterLoginPreCheck = nil }) + + _, err = s.Login(userCtx(), &proto.LoginRequest{ + ProfileName: &target, + Username: &username, + ManagementUrl: "https://mgmt.attacker.example:443", + }) + require.Error(t, err) + require.Equal(t, codes.PermissionDenied, gstatus.Code(err), "want a privilege refusal, got %v", err) + require.False(t, cancelled, "the refused login cancelled the login already in progress") + + active, err := s.profileManager.GetActiveProfileState() + require.NoError(t, err) + require.Equal(t, profilemanager.ID(activeProfile), active.ID, "the refused login switched the active profile anyway") + + stored, err := profilemanager.ReadConfig(targetPath) + require.NoError(t, err) + require.Equal(t, "https://api.netbird.io:443", stored.ManagementURL.String(), "the refused login moved the management URL") +} + +// Login cancels whatever login is already in progress before starting its own. A +// refused caller must not get that far, otherwise anyone able to reach the socket +// can abort someone else's login by sending a request that is denied. +func TestLogin_RefusedChangeLeavesAnInProgressLoginAlone(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + target := "ssh-enabled" + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: filepath.Join(profilemanager.DefaultConfigPathDir, target+".json"), + ManagementURL: "https://api.netbird.io:443", + ServerSSHAllowed: boolPtr(true), + }) + require.NoError(t, err) + + cancelled := false + s.actCancel = func() { cancelled = true } + + _, err = s.Login(userCtx(), &proto.LoginRequest{ + ProfileName: &target, + Username: &username, + ManagementUrl: "https://mgmt.attacker.example:443", + }) + require.Error(t, err) + require.Equal(t, codes.PermissionDenied, gstatus.Code(err), "want a privilege refusal, got %v", err) + require.False(t, cancelled, "the refused login cancelled the login already in progress") +} diff --git a/client/server/login_outcome_test.go b/client/server/login_outcome_test.go new file mode 100644 index 000000000..7ebf04f92 --- /dev/null +++ b/client/server/login_outcome_test.go @@ -0,0 +1,110 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "os" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/proto" +) + +// A login that never reached Management is not a decision about the peer's +// credentials, so it must come back as a retryable error rather than an SSO +// prompt: the user cannot finish a browser login while Management is down, and +// the CLI's own backoff resolves the outage on its own once the daemon reports +// the failure. Reproduces `netbird down; netbird up` printing a device-code URL +// because Management happened to be restarting when the daemon dialed it. +func TestLogin_ManagementUnreachableIsReturnedInsteadOfDemandingSSO(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + unreachable := errors.New("create connection: dial context: context deadline exceeded") + attempts := 0 + s.isLoginRequiredFn = func(context.Context) (bool, error) { + attempts++ + return false, unreachable + } + + resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) + require.Error(t, err) + require.ErrorIs(t, err, unreachable, "the transport failure was replaced by something else") + require.Nil(t, resp, "a failed login must not answer with a login response") + require.Equal(t, 1, attempts) + require.Nil(t, s.oauthAuthFlow.flow, "the daemon started an SSO flow for a peer whose login was never decided") + + status, err := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, err) + require.Equal(t, internal.StatusLoginFailed, status, + "a peer that could not reach Management is not waiting on a login") +} + +// The counterpart: Management refusing the peer's credentials is a decision, and +// the SSO flow still has to start for it. The profile carries an unusable +// private key so the flow setup fails immediately instead of dialing, which is +// enough to show the branch was entered — the refusal itself is never what comes +// back out. +func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) { + s, _, _, username, cfgPath := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + breakProfilePrivateKey(t, cfgPath) + + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return true, nil + } + + _, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) + require.Error(t, err) + + status, stateErr := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, stateErr) + require.Equal(t, internal.StatusLoginFailed, status, + "the SSO flow setup was never reached with the broken key") +} + +func TestLogin_SetupKeyStillRunsWhenPeerNeedsLogin(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return true, nil + } + + var keysTried []string + s.loginAttemptFn = func(_ context.Context, setupKey, _ string) (internal.StatusType, error) { + keysTried = append(keysTried, setupKey) + return "", nil + } + + setupKey := "A2C8E32F-AEB2-4B45-8FD3-8A0C1B2D3E4F" + resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username, SetupKey: setupKey}) + require.NoError(t, err, "the probe's outcome leaked out as the login result") + require.NotNil(t, resp) + require.Equal(t, []string{setupKey}, keysTried, "the setup key never reached the login attempt") + require.Nil(t, s.oauthAuthFlow.flow, "a setup-key login started an SSO flow") + + status, err := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, err) + require.Equal(t, internal.StatusIdle, status) +} + +// breakProfilePrivateKey replaces the profile's private key with an unparseable +// one, which makes any attempt to build a Management client fail on the spot. +func breakProfilePrivateKey(t *testing.T, cfgPath string) { + t.Helper() + + raw, err := os.ReadFile(cfgPath) + require.NoError(t, err) + + var cfg map[string]any + require.NoError(t, json.Unmarshal(raw, &cfg)) + cfg["PrivateKey"] = "not-a-key" + + patched, err := json.Marshal(cfg) + require.NoError(t, err) + require.NoError(t, os.WriteFile(cfgPath, patched, 0o600)) +} diff --git a/client/server/network.go b/client/server/network.go index c38715256..c390b8180 100644 --- a/client/server/network.go +++ b/client/server/network.go @@ -8,7 +8,6 @@ import ( "sort" "strings" - "golang.org/x/exp/maps" "google.golang.org/grpc/codes" gstatus "google.golang.org/grpc/status" @@ -161,30 +160,11 @@ func (s *Server) SelectNetworks(_ context.Context, req *proto.SelectNetworksRequ return nil, fmt.Errorf("no route manager") } - routeSelector := routeManager.GetRouteSelector() if req.GetAll() { - routeSelector.SelectAllRoutes() - } else { - routes := toNetIDs(req.GetNetworkIDs()) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - netIdRoutes := maps.Keys(routesMap) - if err := routeSelector.SelectRoutes(routes, req.GetAppend(), netIdRoutes); err != nil { - return nil, fmt.Errorf("select routes: %w", err) - } - - // Exit nodes are mutually exclusive: if this selection activates an - // exit node, deselect every other available exit node so two can't be - // selected at once. Non-exit route selections are left untouched. - if requestActivatesExitNode(routes, routesMap) { - if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 { - if err := routeSelector.DeselectRoutes(others, netIdRoutes); err != nil { - return nil, fmt.Errorf("deselect sibling exit nodes: %w", err) - } - } - } + routeManager.SelectAllRoutes() + } else if err := routeManager.SelectRoutes(toNetIDs(req.GetNetworkIDs()), req.GetAppend()); err != nil { + return nil, err } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) s.statusRecorder.PublishEvent( proto.SystemEvent_INFO, @@ -224,19 +204,11 @@ func (s *Server) DeselectNetworks(_ context.Context, req *proto.SelectNetworksRe return nil, fmt.Errorf("no route manager") } - routeSelector := routeManager.GetRouteSelector() if req.GetAll() { - routeSelector.DeselectAllRoutes() - } else { - routes := toNetIDs(req.GetNetworkIDs()) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - netIdRoutes := maps.Keys(routesMap) - if err := routeSelector.DeselectRoutes(routes, netIdRoutes); err != nil { - return nil, fmt.Errorf("deselect routes: %w", err) - } + routeManager.DeselectAllRoutes() + } else if err := routeManager.DeselectRoutes(toNetIDs(req.GetNetworkIDs())); err != nil { + return nil, err } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) s.statusRecorder.PublishEvent( proto.SystemEvent_INFO, @@ -261,37 +233,3 @@ func toNetIDs(routes []string) []route.NetID { return netIDs } -func isExitNodeRoutes(routes []*route.Route) bool { - return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network)) -} - -// requestActivatesExitNode reports whether any requested NetID maps to an exit -// node (default route) in the current route table. -func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool { - for _, id := range requested { - if isExitNodeRoutes(routesMap[id]) { - return true - } - } - return false -} - -// otherExitNodeIDs returns every available exit-node NetID that is not in the -// requested set — the siblings to deselect so a single exit node stays active. -func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID { - keep := make(map[route.NetID]struct{}, len(requested)) - for _, id := range requested { - keep[id] = struct{}{} - } - var others []route.NetID - for id, routes := range routesMap { - if !isExitNodeRoutes(routes) { - continue - } - if _, ok := keep[id]; ok { - continue - } - others = append(others, id) - } - return others -} diff --git a/client/server/network_exitnode_test.go b/client/server/network_exitnode_test.go deleted file mode 100644 index 1c0ba0ecb..000000000 --- a/client/server/network_exitnode_test.go +++ /dev/null @@ -1,26 +0,0 @@ -package server - -import ( - "net/netip" - "testing" - - "github.com/stretchr/testify/assert" - - "github.com/netbirdio/netbird/route" -) - -func TestExitNodeSelectionHelpers(t *testing.T) { - routesMap := map[route.NetID][]*route.Route{ - "exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}}, - "exitB": {{Network: netip.MustParsePrefix("::/0")}}, - "lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}}, - } - - assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node") - assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node") - assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node") - assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node") - - others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"}) - assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored") -} diff --git a/client/server/panic_windows.go b/client/server/panic_windows.go index 8592f12ad..4bed6662f 100644 --- a/client/server/panic_windows.go +++ b/client/server/panic_windows.go @@ -3,6 +3,7 @@ package server import ( + "errors" "fmt" "os" "path" @@ -69,7 +70,7 @@ func setStdHandle(f *os.File) error { handle := f.Fd() r0, _, e1 := setStdHandleFn.Call(stdErrorHandle, handle) if r0 == 0 { - if e1 != nil { + if !errors.Is(e1, syscall.Errno(0)) { return e1 } return syscall.EINVAL diff --git a/client/server/server.go b/client/server/server.go index 2b919d58d..f33e19075 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -72,6 +72,9 @@ type Server struct { // RegisterUILog. Guarded by mutex. Consumed by DebugBundle so the bundle // can collect the GUI log even though the daemon runs as root and can't // resolve the user's config dir. Last-writer-wins (one UI per socket). + // DebugBundle opens it on behalf of the bundle requester and refuses a file + // that caller does not own, so a local user cannot read another user's log + // or a root-only file through it. uiLogPath string oauthAuthFlow oauthAuthFlow @@ -82,6 +85,12 @@ type Server struct { // extend flow or vice versa. extendAuthSessionFlow *auth.PendingFlow + // guardedConfigMu serializes a privilege check against the write it + // authorizes. Without it the two are separate steps over the same file, and a + // change that was allowed because the profile had the SSH server disabled + // could land after a concurrent privileged request enabled it. + guardedConfigMu sync.Mutex + mutex sync.Mutex config *profilemanager.Config proto.UnimplementedDaemonServiceServer @@ -126,6 +135,13 @@ type Server struct { updateManager *updater.Manager jwtCache *jwtCache + + // loginAttemptFn stands in for the Management login round trip. Tests set + // it to drive the login outcomes that need a server on the other end; + // production leaves it nil, and every login goes through loginAttempt. + loginAttemptFn func(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) + + isLoginRequiredFn func(ctx context.Context) (bool, error) } type oauthAuthFlow struct { @@ -361,7 +377,34 @@ func (s *Server) connectionGoroutineRunning() bool { } } -// loginAttempt attempts to login using the provided information. it returns a status in case something fails +// attemptLogin runs a login round trip against Management, or the stand-in a +// test installed in place of it. +func (s *Server) attemptLogin(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) { + if s.loginAttemptFn != nil { + return s.loginAttemptFn(ctx, setupKey, jwtToken) + } + return s.loginAttempt(ctx, setupKey, jwtToken) +} + +func (s *Server) isLoginRequired(ctx context.Context) (bool, error) { + if s.isLoginRequiredFn != nil { + return s.isLoginRequiredFn(ctx) + } + + authClient, err := auth.NewAuth(ctx, s.config.PrivateKey, s.config.ManagementURL, s.config) + if err != nil { + log.Errorf("failed to create auth client: %v", err) + return false, err + } + defer authClient.Close() + + return authClient.IsLoginRequired(ctx) +} + +// loginAttempt attempts to login using the provided information. It returns +// StatusNeedsLogin when Management refused the peer's credentials and +// StatusLoginFailed for every other failure, so callers can tell an +// authentication decision apart from a login that never got made. func (s *Server) loginAttempt(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) { authClient, err := auth.NewAuth(ctx, s.config.PrivateKey, s.config.ManagementURL, s.config) if err != nil { @@ -387,6 +430,16 @@ func (s *Server) loginAttempt(ctx context.Context, setupKey, jwtToken string) (i // Login uses setup key to prepare configuration for the daemon. func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigRequest) (*proto.SetConfigResponse, error) { + // Privilege gate: refuse the parts of the request that would let a local + // user turn the root daemon into a root shell. Held across the write so the + // config cannot gain the SSH server between the decision and the update. + // + // Taken before s.mutex: authorizeAndPrepareLogin takes s.mutex while holding + // guardedConfigMu, so acquiring the two in the other order here would let a + // concurrent login deadlock the daemon. + s.guardedConfigMu.Lock() + defer s.guardedConfigMu.Unlock() + s.mutex.Lock() defer s.mutex.Unlock() @@ -411,6 +464,14 @@ func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigReques return nil, err } + stored, err := s.storedProfileConfig(msg.ProfileName, msg.Username) + if err != nil { + return nil, err + } + if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromSetConfig(msg)); err != nil { + return nil, err + } + config, err := s.setConfigInputFromRequest(msg) if err != nil { return nil, err @@ -537,22 +598,23 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro } } - s.mutex.Lock() - if s.actCancel != nil { - s.actCancel() - } - ctx, cancel := context.WithCancel(callerCtx) - - md, ok := metadata.FromIncomingContext(callerCtx) - if ok { - ctx = metadata.NewOutgoingContext(ctx, md) + activeProf, err := s.profileManager.GetActiveProfileState() + if err != nil { + log.Errorf("failed to get active profile state: %v", err) + return nil, fmt.Errorf("failed to get active profile state: %w", err) } - s.actCancel = cancel - s.mutex.Unlock() - - if err := RestoreResidualState(s.rootCtx, s.profileManager.GetStatePath()); err != nil { - log.Warnf(errRestoreResidualState, err) + // Privilege gate: same restrictions as SetConfig, since LoginRequest can carry + // the same fields. It runs before anything here changes daemon state, so a + // refused login neither switches the profile nor cancels a login already in + // progress, and it reads the profile the request targets, which is the one the + // switch below would activate. + stored, err := s.storedLoginConfig(activeProf, msg) + if err != nil { + return nil, err + } + if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil { + return nil, err } state := internal.CtxGetState(s.rootCtx) @@ -563,23 +625,16 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro } }() - activeProf, err := s.profileManager.GetActiveProfileState() + ctx, activeProf, err := s.authorizeAndPrepareLogin(callerCtx, msg, activeProf) if err != nil { - log.Errorf("failed to get active profile state: %v", err) - return nil, fmt.Errorf("failed to get active profile state: %w", err) - } - - if msg.ProfileName != nil { - if _, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf); err != nil { - log.Errorf("failed to switch profile: %v", err) - return nil, err + // The RPC boundary is where this gets recorded: nothing logs handler + // errors for us, and a caller that retries would otherwise leave no + // trace in the daemon log. A refusal is skipped because the gate has + // already logged the decision, with the caller's identity. + if gstatus.Code(err) != codes.PermissionDenied { + log.Errorf("failed to prepare login: %v", err) } - } - - activeProf, err = s.profileManager.GetActiveProfileState() - if err != nil { - log.Errorf("failed to get active profile state: %v", err) - return nil, fmt.Errorf("failed to get active profile state: %w", err) + return nil, err } log.Infof("active profile: %s for %s", activeProf.ID, activeProf.Username) @@ -593,11 +648,6 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro s.mutex.Unlock() - if err := persistLoginOverrides(activeProf, msg.ManagementUrl, msg.OptionalPreSharedKey); err != nil { - log.Errorf("failed to persist login overrides: %v", err) - return nil, fmt.Errorf("persist login overrides: %w", err) - } - config, _, err := s.getConfig(activeProf) if err != nil { log.Errorf("failed to get active profile config: %v", err) @@ -607,7 +657,19 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro s.config = config s.mutex.Unlock() - if _, err := s.loginAttempt(ctx, "", ""); err == nil { + // A probe that errors leaves the login undecided: Management unreachable, a + // restart mid-request, an internal error. Those are returned for the caller + // to retry, because turning them into an SSO prompt asks the user to solve + // something that is not theirs to solve, and a browser login cannot succeed + // while Management is unreachable anyway. Only Management refusing the + // peer's key is a decision, and IsLoginRequired reports that as + // needsLogin=true rather than an error. + needsLogin, err := s.isLoginRequired(ctx) + if err != nil { + state.Set(internal.StatusLoginFailed) + return nil, err + } + if !needsLogin { state.Set(internal.StatusIdle) return &proto.LoginResponse{}, nil } @@ -668,7 +730,7 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro // which returns NeedsLogin and parks on the browser leg. state.Set(internal.StatusConnecting) - if loginStatus, err := s.loginAttempt(ctx, msg.SetupKey, ""); err != nil { + if loginStatus, err := s.attemptLogin(ctx, msg.SetupKey, ""); err != nil { state.Set(loginStatus) return nil, err } @@ -823,7 +885,7 @@ func (s *Server) WaitSSOLogin(callerCtx context.Context, msg *proto.WaitSSOLogin s.oauthAuthFlow.expiresAt = time.Now() s.mutex.Unlock() - if loginStatus, err := s.loginAttempt(ctx, "", tokenInfo.GetTokenToUse()); err != nil { + if loginStatus, err := s.attemptLogin(ctx, "", tokenInfo.GetTokenToUse()); err != nil { state.Set(loginStatus) return nil, err } @@ -980,6 +1042,63 @@ func (s *Server) waitForUp(callerCtx context.Context) (*proto.UpResponse, error) } } +// storedProfileConfig loads the on-disk config of the profile a request +// targets, so a privileged-change decision can be made against the values the +// profile currently holds. A profile that has no config file yet yields nil, +// which every caller must read as "nothing enabled yet". +func (s *Server) storedProfileConfig(handle, username string) (*profilemanager.Config, error) { + resolved, err := s.resolveProfileHandle(handle, username) + if err != nil { + return nil, err + } + + path := resolved.Path + if path == "" { + path = profilemanager.DefaultConfigPath + } + + return s.storedConfigAtPath(path) +} + +// storedLoginConfig loads the on-disk config of the profile a login request +// targets: the one it names, or the active one when it names none. Used to decide +// a privileged change before the request is allowed to switch profiles. +func (s *Server) storedLoginConfig(activeProf *profilemanager.ActiveProfileState, msg *proto.LoginRequest) (*profilemanager.Config, error) { + if msg.ProfileName == nil { + cfgPath, err := activeProf.FilePath() + if err != nil { + return nil, fmt.Errorf("active profile file path: %w", err) + } + return s.storedConfigAtPath(cfgPath) + } + + // Mirrors switchProfileIfNeeded: the default profile resolves without a + // username, so this reads the same profile the switch would activate. + handle := *msg.ProfileName + username := "" + if handle != profilemanager.DefaultProfileName { + username = msg.GetUsername() + } + return s.storedProfileConfig(handle, username) +} + +// storedConfigAtPath reads a profile config file, yielding nil when it does not +// exist yet. +func (s *Server) storedConfigAtPath(path string) (*profilemanager.Config, error) { + if _, err := os.Stat(path); err != nil { + if os.IsNotExist(err) { + return nil, nil //nolint:nilnil + } + return nil, fmt.Errorf("stat profile config: %w", err) + } + + cfg, err := profilemanager.GetConfig(path) + if err != nil { + return nil, fmt.Errorf("read profile config: %w", err) + } + return cfg, nil +} + // resolveProfileHandle resolves a wire-level profile handle (display // name, ID, or unique ID prefix) to a concrete profile. Returns gRPC // status errors so handlers can return them directly. @@ -1081,7 +1200,10 @@ func (s *Server) Down(ctx context.Context, _ *proto.DownRequest) (*proto.DownRes if err := s.cleanupConnection(); err != nil { s.mutex.Unlock() - // todo review to update the status in case any type of error + if errors.Is(err, ErrServiceNotUp) { + log.Debugf("Down called while service not up: %v", err) + return nil, err + } log.Errorf("failed to shut down properly: %v", err) return nil, err } @@ -1154,7 +1276,7 @@ func (s *Server) cleanupConnection() error { // making the run loop the sole owner of engine shutdown. if engine != nil { if err := engine.Stop(); err != nil { - return err + log.Errorf("failed to stop engine during cleanup: %v", err) } } @@ -1194,6 +1316,12 @@ func (s *Server) handleProfileLogout(ctx context.Context, msg *proto.LogoutReque if err := s.logoutFromProfile(ctx, resolved); err != nil { log.Errorf("failed to logout from profile %s: %v", resolved.ID, err) + // A refused deregistration is already a status error carrying the reason + // and the command to run; rewrapping it as Internal would flatten both + // into a gRPC dump for the user. + if _, isStatus := gstatus.FromError(err); isStatus { + return nil, err + } return nil, gstatus.Errorf(codes.Internal, "logout: %v", err) } @@ -1315,6 +1443,13 @@ func (s *Server) sendLogoutRequest(ctx context.Context) error { } func (s *Server) sendLogoutRequestWithConfig(ctx context.Context, config *profilemanager.Config) error { + // Privilege gate: deregistering frees this machine's key to be registered + // against another management server, which is only restricted while the SSH + // server makes that a privilege handover. + if err := requirePrivilegeForDeregistration(ctx, config); err != nil { + return err + } + key, err := wgtypes.ParseKey(config.PrivateKey) if err != nil { return fmt.Errorf("parse private key: %w", err) @@ -1588,8 +1723,8 @@ func (s *Server) RequestJWTAuth( hint = profilemanager.GetLoginHint() } - isDesktop := isUnixRunningDesktop() - oAuthFlow, err := auth.NewOAuthFlow(ctx, config, isDesktop, false, hint) + // the daemon has no graphical session of its own, only the caller can answer this + oAuthFlow, err := auth.NewOAuthFlow(ctx, config, msg.GetHasGraphicalSession(), false, hint) if err != nil { return nil, gstatus.Errorf(codes.Internal, "failed to create OAuth flow: %v", err) } @@ -1680,6 +1815,9 @@ func (s *Server) RequestExtendAuthSession( if connectClient == nil { return nil, gstatus.Errorf(codes.FailedPrecondition, "client is not running") } + if connectClient.Engine() == nil { + return nil, gstatus.Errorf(codes.FailedPrecondition, "session can no longer be extended, log in again to reconnect") + } hint := "" if msg.Hint != nil { @@ -1689,8 +1827,8 @@ func (s *Server) RequestExtendAuthSession( hint = profilemanager.GetLoginHint() } - isDesktop := isUnixRunningDesktop() - oAuthFlow, err := auth.NewOAuthFlow(ctx, config, isDesktop, false, hint) + // the daemon has no graphical session of its own, only the caller can answer this + oAuthFlow, err := auth.NewOAuthFlow(ctx, config, msg.GetHasGraphicalSession(), false, hint) if err != nil { return nil, gstatus.Errorf(codes.Internal, "failed to create OAuth flow: %v", err) } @@ -1862,13 +2000,6 @@ func (s *Server) ExposeService(req *proto.ExposeServiceRequest, srv proto.Daemon return nil } -func isUnixRunningDesktop() bool { - if runtime.GOOS != "linux" && runtime.GOOS != "freebsd" { - return false - } - return os.Getenv("DESKTOP_SESSION") != "" || os.Getenv("XDG_CURRENT_DESKTOP") != "" -} - func (s *Server) runProbes(ctx context.Context, waitForProbeResult bool) { if s.connectClient == nil { return @@ -2060,7 +2191,10 @@ func (s *Server) RemoveProfile(ctx context.Context, msg *proto.RemoveProfileRequ } if err := s.logoutFromProfile(ctx, resolved); err != nil { - log.Warnf("failed to logout from profile %s before removal: %v", resolved.ID, err) + // Deregistration is best-effort here: the local profile is removed + // either way, so an unprivileged caller leaves the peer registered on + // the management server rather than being blocked from removing it. + log.Warnf("removing profile %s locally without deregistering it: %v", resolved.ID, err) } if err := s.profileManager.RemoveProfile(resolved.ID, msg.Username); err != nil { @@ -2357,6 +2491,69 @@ func sendTerminalNotification() error { // persistLoginOverrides writes management URL and pre-shared key from a LoginRequest to the // active profile config so that subsequent reads pick them up. Empty/nil values are ignored. +// afterLoginPreCheck is a seam for tests to run a concurrent config change +// between Login's first privilege check and the authoritative one. +var afterLoginPreCheck func() + +// authorizeAndPrepareLogin makes the authoritative privilege decision for a login +// and, when it passes, carries out every state change that decision authorizes: +// cancelling an login already in progress, switching to the requested profile, and +// persisting the config overrides the request carries. +// +// All of it happens under guardedConfigMu, which SetConfig also holds across its +// own check and write. Login's earlier check refuses the ordinary case before any +// of this is reached; this one exists because that check is not synchronized +// against a concurrent privileged request that enables the SSH server, and a +// caller refused here must not have cancelled or switched anything either. +func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto.LoginRequest, activeProf *profilemanager.ActiveProfileState) (context.Context, *profilemanager.ActiveProfileState, error) { + if afterLoginPreCheck != nil { + afterLoginPreCheck() + } + + s.guardedConfigMu.Lock() + defer s.guardedConfigMu.Unlock() + + stored, err := s.storedLoginConfig(activeProf, msg) + if err != nil { + return nil, nil, err + } + if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil { + return nil, nil, err + } + + s.mutex.Lock() + if s.actCancel != nil { + s.actCancel() + } + ctx, cancel := context.WithCancel(callerCtx) + if md, ok := metadata.FromIncomingContext(callerCtx); ok { + ctx = metadata.NewOutgoingContext(ctx, md) + } + s.actCancel = cancel + s.mutex.Unlock() + + if err := RestoreResidualState(s.rootCtx, s.profileManager.GetStatePath()); err != nil { + log.Warnf(errRestoreResidualState, err) + } + + if msg.ProfileName != nil { + if _, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf); err != nil { + return nil, nil, fmt.Errorf("switch profile: %w", err) + } + } + + activeProf, err = s.profileManager.GetActiveProfileState() + if err != nil { + return nil, nil, fmt.Errorf("active profile state: %w", err) + } + + if err := persistLoginOverrides(activeProf, msg.ManagementUrl, msg.OptionalPreSharedKey); err != nil { + return nil, nil, fmt.Errorf("persist login overrides: %w", err) + } + + return ctx, activeProf, nil +} + func persistLoginOverrides(activeProf *profilemanager.ActiveProfileState, managementURL string, preSharedKey *string) error { if preSharedKey != nil && *preSharedKey == "" { preSharedKey = nil diff --git a/client/server/server_privileged_test.go b/client/server/server_privileged_test.go index 8b6f78f04..0366ccb31 100644 --- a/client/server/server_privileged_test.go +++ b/client/server/server_privileged_test.go @@ -10,7 +10,7 @@ import ( "testing" "time" - "github.com/golang/mock/gomock" + "go.uber.org/mock/gomock" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel" diff --git a/client/server/setconfig_mdm_test.go b/client/server/setconfig_mdm_test.go index 9baf16136..ae323ea8c 100644 --- a/client/server/setconfig_mdm_test.go +++ b/client/server/setconfig_mdm_test.go @@ -66,7 +66,11 @@ func setupServerWithProfile(t *testing.T) (s *Server, ctx context.Context, profN Username: currUser.Username, })) - ctx = context.Background() + // The privileged-change gate reads the caller's kernel identity from the + // context, which a real caller gets from the daemon's transport credentials. + // This test drives the handler directly, so it stands in for a root caller; + // without an identity the gate would (correctly) refuse the SSH fields. + ctx = privilegedTestCtx() s = New(ctx, "console", "", false, false, false, false) return s, ctx, profName, currUser.Username, cfgPath } diff --git a/client/server/setconfig_test.go b/client/server/setconfig_test.go index 0e55257a9..db7a26f03 100644 --- a/client/server/setconfig_test.go +++ b/client/server/setconfig_test.go @@ -1,7 +1,6 @@ package server import ( - "context" "os/user" "path/filepath" "reflect" @@ -52,7 +51,11 @@ func TestSetConfig_AllFieldsSaved(t *testing.T) { }) require.NoError(t, err) - ctx := context.Background() + // The privileged-change gate reads the caller's kernel identity from the + // context, which a real caller gets from the daemon's transport credentials. + // This test drives the handler directly, so it stands in for a root caller; + // without an identity the gate would (correctly) refuse the SSH fields. + ctx := privilegedTestCtx() s := New(ctx, "console", "", false, false, false, false) rosenpassEnabled := true diff --git a/client/server/ssh_gate.go b/client/server/ssh_gate.go new file mode 100644 index 000000000..ca1b4c4ee --- /dev/null +++ b/client/server/ssh_gate.go @@ -0,0 +1,282 @@ +package server + +import ( + "context" + "fmt" + "net/url" + "runtime" + "strings" + + log "github.com/sirupsen/logrus" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/proto" + "github.com/netbirdio/netbird/util" +) + +// The daemon runs as root/LocalSystem, so a handful of config changes cross the +// user-to-root boundary and are restricted to privileged callers: +// +// - Enabling SSH root login, or disabling SSH authentication, turns the +// daemon's SSH server into a root (or unauthenticated) shell. +// - Enabling the SSH server at all is what makes the above reachable, and a +// profile the caller owns is not a privilege they hold. +// - While the SSH server is enabled, repointing the profile at another +// management identity hands SSH authorization decisions, including which +// keys and users are accepted, to whoever controls that identity. Changing +// the management URL and deregistering the peer are both ways to do that. +// +// Everything else stays unauthenticated, so this is not an authorization model: +// it only refuses the changes that would let a local user become root. A caller +// whose identity cannot be established is refused as well. + +// privilegedConfigChange is the subset of a config request that crosses the +// user-to-root boundary. Fields are nil or empty when the request leaves them +// untouched. +type privilegedConfigChange struct { + managementURL string + serverSSHAllowed *bool + enableSSHRoot *bool + disableSSHAuth *bool +} + +func privilegedChangeFromSetConfig(msg *proto.SetConfigRequest) privilegedConfigChange { + return privilegedConfigChange{ + managementURL: msg.GetManagementUrl(), + serverSSHAllowed: msg.ServerSSHAllowed, + enableSSHRoot: msg.EnableSSHRoot, + disableSSHAuth: msg.DisableSSHAuth, + } +} + +func privilegedChangeFromLogin(msg *proto.LoginRequest) privilegedConfigChange { + return privilegedConfigChange{ + managementURL: msg.GetManagementUrl(), + serverSSHAllowed: msg.ServerSSHAllowed, + enableSSHRoot: msg.EnableSSHRoot, + disableSSHAuth: msg.DisableSSHAuth, + } +} + +// requirePrivilegeForConfigChange refuses the privileged parts of a config +// change when the caller is not root/administrator. stored is the profile's +// current config, or nil when it has none yet. +// +// Each check compares against the stored value so that a request restating a +// value it does not change is never refused: a UI that submits the whole +// settings form must not start failing once an administrator has enabled SSH. +func requirePrivilegeForConfigChange(ctx context.Context, stored *profilemanager.Config, change privilegedConfigChange) error { + if enables(storedFlag(stored, func(c *profilemanager.Config) *bool { return c.EnableSSHRoot }), change.enableSSHRoot) { + return denyPrivileged(ctx, "enabling SSH root login", ipcauth.UpCommand("--enable-ssh-root")) + } + + if enables(storedFlag(stored, func(c *profilemanager.Config) *bool { return c.DisableSSHAuth }), change.disableSSHAuth) { + return denyPrivileged(ctx, "disabling SSH authentication", ipcauth.UpCommand("--disable-ssh-auth")) + } + + if enables(sshServerCurrentlyAllowed(stored), change.serverSSHAllowed) { + return denyPrivileged(ctx, "enabling the NetBird SSH server", ipcauth.UpCommand("--allow-server-ssh")) + } + + // Only guard the management binding while the SSH server is enabled: that is + // when the management identity decides who may open a shell here. + if !sshServerEnabled(stored) { + return nil + } + + if change.managementURL != "" && !sameManagementURL(stored.ManagementURL, change.managementURL) { + return denyPrivileged(ctx, + "changing the management URL while the NetBird SSH server is enabled", + ipcauth.UpCommand("-m "+change.managementURL)) + } + + return nil +} + +// requirePrivilegeForDeregistration refuses to deregister the peer from the +// management server when the caller is not privileged and the profile has the +// SSH server enabled. Deregistering frees the peer's key to be registered +// against another management identity, which is the same handover the +// management URL check refuses. +// +// Callers that treat deregistration as best-effort (profile removal) continue +// without it; callers that were asked to deregister surface the error. +func requirePrivilegeForDeregistration(ctx context.Context, cfg *profilemanager.Config) error { + if !sshServerEnabled(cfg) { + return nil + } + + return denyPrivileged(ctx, + "deregistering this peer while the NetBird SSH server is enabled", + ipcauth.ElevatedCommand("netbird logout")) +} + +// denyPrivileged returns nil when the caller is privileged, and otherwise a +// PermissionDenied whose message names the action and the command that performs +// it with the privileges it needs. The same summary and command ride along as an +// ErrorInfo detail so the CLI and the UI can present them without parsing text. +// +// action reads as the subject of a sentence ("enabling SSH root login"), and +// command is the equivalent command, already elevated for the platform. +func denyPrivileged(ctx context.Context, action, command string) error { + id, ok := ipcauth.CallerIdentity(ctx) + if !ok { + log.Warnf("denying %s: the caller's identity cannot be verified on this control channel", action) + return privilegeError(unidentifiedSummary(action), reinstallCommand()) + } + + if ipcauth.IsPrivilegedCaller(id) { + log.Infof("allowing %s for privileged caller %s", action, id) + return nil + } + + log.Warnf("denying %s for unprivileged caller %s", action, id) + actor, command := requiredActor(command) + return privilegeError(privilegeSummary(action, actor), command) +} + +// requiredActor names who may perform the operation and adjusts the command to +// match. A daemon that is not itself privileged delegates to its own identity, so +// telling that host's user to become root is wrong twice over: root is not what the +// daemon checks for, and a rootless container has neither root nor sudo. +func requiredActor(command string) (string, string) { + self, delegates := ipcauth.SelfDelegatesTo() + if !delegates { + return ipcauth.PrivilegedActor(), command + } + return fmt.Sprintf("the user the daemon runs as (%s)", self), strings.ReplaceAll(command, "sudo ", "") +} + +// privilegeError builds the PermissionDenied carrying summary and command. +func privilegeError(summary, command string) error { + st := gstatus.New(codes.PermissionDenied, fmt.Sprintf("%s\n\n%s", summary, command)) + + detailed, err := st.WithDetails(&errdetails.ErrorInfo{ + Reason: ipcauth.ErrorReasonPrivilegeRequired, + Domain: ipcauth.ErrorDomain, + Metadata: map[string]string{ + ipcauth.ErrorMetaSummary: summary, + ipcauth.ErrorMetaCommand: command, + }, + }) + if err != nil { + log.Debugf("attach privilege error detail: %v", err) + return st.Err() + } + return detailed.Err() +} + +// privilegeSummary states what is refused and what it needs, in one sentence +// that reads the same in a dialog and in a terminal. +func privilegeSummary(action, actor string) string { + return fmt.Sprintf("%s requires %s.", capitalize(action), actor) +} + +// unidentifiedSummary covers a control channel that carries no caller identity. +// Elevating does not help there, so it points at the daemon's socket instead. +func unidentifiedSummary(action string) string { + return fmt.Sprintf("%s requires %s, and the daemon cannot verify who is calling over its current socket. "+ + "Reinstall the service on a socket that carries the caller's identity.", capitalize(action), ipcauth.PrivilegedActor()) +} + +// reinstallCommand is the command that moves the daemon onto a socket whose +// callers can be identified. +func reinstallCommand() string { + if runtime.GOOS == "windows" { + return fmt.Sprintf("netbird service install --daemon-addr %s", daemonaddr.WindowsPipeAddr) + } + return "sudo netbird service install --daemon-addr unix:///var/run/netbird.sock" +} + +func capitalize(s string) string { + if s == "" { + return s + } + return strings.ToUpper(s[:1]) + s[1:] +} + +// enables reports whether requested turns a flag on that is currently off. A +// request that restates the stored value, or turns the flag off, is not a +// privileged change. +func enables(stored, requested *bool) bool { + if requested == nil || !*requested { + return false + } + return stored == nil || !*stored +} + +// storedFlag reads a flag from the stored config, tolerating a config that does +// not exist yet. +func storedFlag(cfg *profilemanager.Config, get func(*profilemanager.Config) *bool) *bool { + if cfg == nil { + return nil + } + return get(cfg) +} + +// sshServerEnabled reports whether the profile currently runs the SSH server. +// +// A nil flag means ON, matching what the engine does with the same config +// (util.ReturnBoolWithDefaultTrue in internal/connect.go, kept for configs written +// before the flag existed). Reading it as OFF here would open the management-URL +// and deregistration guards on exactly those legacy hosts, whose SSH server is +// running. Configs loaded through profilemanager have already been materialised by +// apply(), so this is the same answer by a route that does not depend on that. +func sshServerEnabled(cfg *profilemanager.Config) bool { + if cfg == nil { + return false + } + return util.ReturnBoolWithDefaultTrue(cfg.ServerSSHAllowed) +} + +// sshServerCurrentlyAllowed is the value an enable request is compared against. It +// shares sshServerEnabled's nil-means-on default, so restating "on" for a legacy +// config is correctly seen as no change. +func sshServerCurrentlyAllowed(cfg *profilemanager.Config) *bool { + enabled := sshServerEnabled(cfg) + if cfg == nil { + return nil + } + return &enabled +} + +// sameManagementURL reports whether requested addresses the same management +// server as stored, comparing scheme, host and effective port so that an +// equivalent spelling ("https://api.netbird.io" for a stored +// "https://api.netbird.io:443") is not treated as a change. It fails closed: +// anything unparseable counts as a change and therefore needs privilege. +func sameManagementURL(stored *url.URL, requested string) bool { + if stored == nil { + return false + } + + // Normalise the requested URL through the config layer's own parser, so the + // comparison cannot drift from how the value would actually be stored. + parsed, err := profilemanager.ParseServiceURL("Management URL", requested) + if err != nil { + return false + } + + return stored.Scheme == parsed.Scheme && + stored.Hostname() == parsed.Hostname() && + effectivePort(stored) == effectivePort(parsed) +} + +func effectivePort(u *url.URL) string { + if port := u.Port(); port != "" { + return port + } + switch u.Scheme { + case "https": + return "443" + case "http": + return "80" + default: + return "" + } +} diff --git a/client/server/ssh_gate_test.go b/client/server/ssh_gate_test.go new file mode 100644 index 000000000..cbd345f16 --- /dev/null +++ b/client/server/ssh_gate_test.go @@ -0,0 +1,348 @@ +package server + +import ( + "context" + "net/url" + "os" + "runtime" + "strings" + "testing" + + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/internal/profilemanager" +) + +// ctxWithIdentity builds a request context carrying the identity the transport +// credentials would have attached. +func ctxWithIdentity(id ipcauth.Identity) context.Context { + return peer.NewContext(context.Background(), &peer.Peer{ + AuthInfo: ipcauth.AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, + }) +} + +// unprivUID is deliberately not this process's own uid. An unprivileged daemon +// treats a caller sharing its identity as privileged (rootless containers), and +// the test binary would otherwise stand in for both the daemon and the caller. +// os.Geteuid returns -1 on Windows, where identities are SIDs instead and this is +// unused. +var unprivUID = uint32(os.Geteuid() + 1) + +// The fabricated identities have to be shaped like the platform's: a uid says +// nothing on Windows, and a zero uid there would read as root and be privileged. +func rootCtx() context.Context { return ctxWithIdentity(privilegedIdentity()) } +func userCtx() context.Context { return ctxWithIdentity(unprivilegedIdentity()) } + +func privilegedIdentity() ipcauth.Identity { + if runtime.GOOS == "windows" { + // LocalSystem, which is what the Windows service account is. + return ipcauth.Identity{SID: "S-1-5-18"} + } + return ipcauth.Identity{UID: 0} +} + +func unprivilegedIdentity() ipcauth.Identity { + if runtime.GOOS == "windows" { + // A plain user SID: no groups, so no BUILTIN\Administrators, and not + // elevated. + return ipcauth.Identity{SID: "S-1-5-21-1-2-3-1001"} + } + return ipcauth.Identity{UID: unprivUID, GID: unprivUID} +} +func noIdentityCtx() context.Context { return context.Background() } + +func boolPtr(v bool) *bool { return &v } + +func mustURL(t *testing.T, raw string) *url.URL { + t.Helper() + u, err := url.Parse(raw) + if err != nil { + t.Fatalf("parse %q: %v", raw, err) + } + return u +} + +func assertDenied(t *testing.T, err error) { + t.Helper() + if err == nil { + t.Fatal("expected the change to be refused, got nil") + } + st := gstatus.Convert(err) + if st.Code() != codes.PermissionDenied { + t.Fatalf("code = %v, want PermissionDenied", st.Code()) + } + // The refusal must be machine-readable: the CLI and the UI render the + // summary and command from the detail rather than parsing the message. + var info *errdetails.ErrorInfo + for _, d := range st.Details() { + if got, ok := d.(*errdetails.ErrorInfo); ok { + info = got + } + } + if info == nil { + t.Fatal("refusal carries no ErrorInfo detail") + } + if info.GetReason() != ipcauth.ErrorReasonPrivilegeRequired || info.GetDomain() != ipcauth.ErrorDomain { + t.Fatalf("detail = %s/%s, want %s/%s", info.GetDomain(), info.GetReason(), ipcauth.ErrorDomain, ipcauth.ErrorReasonPrivilegeRequired) + } + if info.GetMetadata()[ipcauth.ErrorMetaSummary] == "" { + t.Error("detail carries no summary") + } + if info.GetMetadata()[ipcauth.ErrorMetaCommand] == "" { + t.Error("detail carries no command") + } +} + +func assertAllowed(t *testing.T, err error) { + t.Helper() + if err != nil { + t.Fatalf("expected the change to be allowed, got %v", err) + } +} + +func TestRequirePrivilegeForConfigChange_SSHFlags(t *testing.T) { + tests := []struct { + name string + stored *profilemanager.Config + change privilegedConfigChange + privileged bool + wantDeny bool + }{ + { + name: "enabling the ssh server unprivileged is refused", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + wantDeny: true, + }, + { + name: "enabling the ssh server as root is allowed", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + privileged: true, + }, + { + name: "restating an already enabled ssh server is not a change", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + }, + { + name: "turning the ssh server off is not guarded", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(false)}, + }, + { + name: "a profile with no config yet counts as off, so enabling is refused", + stored: nil, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + wantDeny: true, + }, + { + name: "enabling ssh root login unprivileged is refused", + stored: &profilemanager.Config{EnableSSHRoot: boolPtr(false)}, + change: privilegedConfigChange{enableSSHRoot: boolPtr(true)}, + wantDeny: true, + }, + { + name: "restating ssh root login is not a change", + stored: &profilemanager.Config{EnableSSHRoot: boolPtr(true)}, + change: privilegedConfigChange{enableSSHRoot: boolPtr(true)}, + }, + { + name: "turning ssh root login off is not guarded", + stored: &profilemanager.Config{EnableSSHRoot: boolPtr(true)}, + change: privilegedConfigChange{enableSSHRoot: boolPtr(false)}, + }, + { + name: "disabling ssh authentication unprivileged is refused", + stored: &profilemanager.Config{DisableSSHAuth: boolPtr(false)}, + change: privilegedConfigChange{disableSSHAuth: boolPtr(true)}, + wantDeny: true, + }, + { + name: "re-enabling ssh authentication is not guarded", + stored: &profilemanager.Config{DisableSSHAuth: boolPtr(true)}, + change: privilegedConfigChange{disableSSHAuth: boolPtr(false)}, + }, + { + name: "a request that touches none of the guarded fields is allowed", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + change: privilegedConfigChange{}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := userCtx() + if tt.privileged { + ctx = rootCtx() + } + err := requirePrivilegeForConfigChange(ctx, tt.stored, tt.change) + if tt.wantDeny { + assertDenied(t, err) + return + } + assertAllowed(t, err) + }) + } +} + +func TestRequirePrivilegeForConfigChange_ManagementURL(t *testing.T) { + sshOn := func(raw string) *profilemanager.Config { + return &profilemanager.Config{ServerSSHAllowed: boolPtr(true), ManagementURL: mustURL(t, raw)} + } + sshOff := func(raw string) *profilemanager.Config { + return &profilemanager.Config{ServerSSHAllowed: boolPtr(false), ManagementURL: mustURL(t, raw)} + } + + tests := []struct { + name string + stored *profilemanager.Config + requested string + privileged bool + wantDeny bool + }{ + { + name: "moving the binding while ssh is enabled is refused", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://attacker.example.com:443", + wantDeny: true, + }, + { + name: "moving the binding as root is allowed", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://selfhosted.example.com:443", + privileged: true, + }, + { + name: "the same url restated is not a change", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://api.netbird.io:443", + }, + { + name: "an equivalent spelling of the same url is not a change", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://api.netbird.io", + }, + { + name: "an equivalent spelling with an explicit http port is not a change", + stored: sshOn("http://mgmt.internal:80"), + requested: "http://mgmt.internal", + }, + { + name: "a different port on the same host is a change", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://api.netbird.io:8443", + wantDeny: true, + }, + { + name: "a different scheme on the same host is a change", + stored: sshOn("https://mgmt.internal:443"), + requested: "http://mgmt.internal:443", + wantDeny: true, + }, + { + name: "with ssh disabled the binding is not guarded at all", + stored: sshOff("https://api.netbird.io:443"), + requested: "https://attacker.example.com:443", + }, + { + name: "an unparseable url fails closed", + stored: sshOn("https://api.netbird.io:443"), + requested: "ht tp://%zz", + wantDeny: true, + }, + { + name: "an empty url leaves the binding alone", + stored: sshOn("https://api.netbird.io:443"), + requested: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := userCtx() + if tt.privileged { + ctx = rootCtx() + } + err := requirePrivilegeForConfigChange(ctx, tt.stored, privilegedConfigChange{managementURL: tt.requested}) + if tt.wantDeny { + assertDenied(t, err) + return + } + assertAllowed(t, err) + }) + } +} + +// A caller the daemon cannot identify must be refused, not trusted: that is the +// state on a TCP daemon socket, where no peer credentials exist. +func TestRequirePrivilegeForConfigChange_UnidentifiedCallerIsRefused(t *testing.T) { + err := requirePrivilegeForConfigChange(noIdentityCtx(), + &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + privilegedConfigChange{serverSSHAllowed: boolPtr(true)}) + assertDenied(t, err) + + // The guidance must point at the socket rather than at sudo, since elevating + // would not help. + st := gstatus.Convert(err) + if !strings.Contains(st.Message(), "service install") { + t.Errorf("message %q does not tell the operator how to fix the socket", st.Message()) + } +} + +func TestRequirePrivilegeForDeregistration(t *testing.T) { + tests := []struct { + name string + cfg *profilemanager.Config + privileged bool + wantDeny bool + }{ + { + name: "deregistering while ssh is enabled is refused", + cfg: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + wantDeny: true, + }, + { + name: "deregistering while ssh is enabled is allowed for root", + cfg: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + privileged: true, + }, + { + name: "deregistering with ssh disabled is not guarded", + cfg: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + }, + { + name: "deregistering a profile with no config is not guarded", + cfg: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := userCtx() + if tt.privileged { + ctx = rootCtx() + } + err := requirePrivilegeForDeregistration(ctx, tt.cfg) + if tt.wantDeny { + assertDenied(t, err) + return + } + assertAllowed(t, err) + }) + } +} + +// privilegedTestCtx is the context a handler-level test should use when it is +// standing in for a root/administrator caller. Tests that drive the handlers +// directly have no transport credentials, and the privileged-change gate refuses +// a caller it cannot identify. +func privilegedTestCtx() context.Context { return rootCtx() } diff --git a/client/ssh/client/client.go b/client/ssh/client/client.go index ebf8eb794..31143a4f4 100644 --- a/client/ssh/client/client.go +++ b/client/ssh/client/client.go @@ -9,7 +9,6 @@ import ( "path/filepath" "runtime" "strconv" - "strings" "time" log "github.com/sirupsen/logrus" @@ -17,7 +16,6 @@ import ( "golang.org/x/crypto/ssh/knownhosts" "golang.org/x/term" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/profilemanager" @@ -32,7 +30,7 @@ const ( // DefaultDaemonAddr is the default address for the NetBird daemon DefaultDaemonAddr = "unix:///var/run/netbird.sock" // DefaultDaemonAddrWindows is the default address for the NetBird daemon on Windows - DefaultDaemonAddrWindows = "tcp://127.0.0.1:41731" + DefaultDaemonAddrWindows = daemonaddr.WindowsPipeAddr ) // Client wraps crypto/ssh Client for simplified SSH operations @@ -268,7 +266,7 @@ func getDefaultDaemonAddr() string { return addr } if runtime.GOOS == "windows" { - return DefaultDaemonAddrWindows + return daemonaddr.ResolveDaemonAddr(DefaultDaemonAddrWindows) } return daemonaddr.ResolveUnixDaemonAddr(DefaultDaemonAddr) } @@ -315,21 +313,23 @@ func Dial(ctx context.Context, addr, user string, opts DialOptions) (*Client, er // dialSSH establishes an SSH connection without JWT authentication func dialSSH(ctx context.Context, network, addr string, config *ssh.ClientConfig) (*Client, error) { + if config.Timeout > 0 { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, config.Timeout) + defer cancel() + } + dialer := &net.Dialer{} conn, err := dialer.DialContext(ctx, network, addr) if err != nil { return nil, fmt.Errorf("dial %s: %w", addr, err) } - clientConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config) + client, err := nbssh.Handshake(ctx, conn, addr, config) if err != nil { - if closeErr := conn.Close(); closeErr != nil { - log.Debugf("connection close after handshake failure: %v", closeErr) - } - return nil, fmt.Errorf("ssh handshake: %w", err) + return nil, err } - client := ssh.NewClient(clientConn, chans, reqs) return &Client{ client: client, }, nil @@ -410,12 +410,9 @@ func verifyHostKeyViaDaemon(hostname string, remote net.Addr, key ssh.PublicKey, } func connectToDaemon(daemonAddr string) (*grpc.ClientConn, error) { - addr := strings.TrimPrefix(daemonAddr, "tcp://") + target, opts := daemonaddr.DialTarget(daemonAddr) - conn, err := grpc.NewClient( - addr, - grpc.WithTransportCredentials(insecure.NewCredentials()), - ) + conn, err := grpc.NewClient(target, opts...) if err != nil { log.Debugf("failed to create gRPC client for NetBird daemon at %s: %v", daemonAddr, err) return nil, fmt.Errorf("failed to connect to NetBird daemon: %w", err) diff --git a/client/ssh/client/terminal_unix.go b/client/ssh/client/terminal_unix.go index aaa3418f9..a963dc8be 100644 --- a/client/ssh/client/terminal_unix.go +++ b/client/ssh/client/terminal_unix.go @@ -12,6 +12,8 @@ import ( log "github.com/sirupsen/logrus" "golang.org/x/crypto/ssh" "golang.org/x/term" + + nbssh "github.com/netbirdio/netbird/client/ssh" ) func (c *Client) setupTerminalMode(ctx context.Context, session *ssh.Session) error { @@ -82,37 +84,7 @@ func (c *Client) setupTerminal(session *ssh.Session, fd int) error { return fmt.Errorf("get terminal size: %w", err) } - modes := ssh.TerminalModes{ - ssh.ECHO: 1, - ssh.TTY_OP_ISPEED: 14400, - ssh.TTY_OP_OSPEED: 14400, - // Ctrl+C - ssh.VINTR: 3, - // Ctrl+\ - ssh.VQUIT: 28, - // Backspace - ssh.VERASE: 127, - // Ctrl+U - ssh.VKILL: 21, - // Ctrl+D - ssh.VEOF: 4, - ssh.VEOL: 0, - ssh.VEOL2: 0, - // Ctrl+Q - ssh.VSTART: 17, - // Ctrl+S - ssh.VSTOP: 19, - // Ctrl+Z - ssh.VSUSP: 26, - // Ctrl+O - ssh.VDISCARD: 15, - // Ctrl+R - ssh.VREPRINT: 18, - // Ctrl+W - ssh.VWERASE: 23, - // Ctrl+V - ssh.VLNEXT: 22, - } + modes := nbssh.DefaultTerminalModes terminal := os.Getenv("TERM") if terminal == "" { diff --git a/client/ssh/client/terminal_windows.go b/client/ssh/client/terminal_windows.go index 462438317..c6156fc26 100644 --- a/client/ssh/client/terminal_windows.go +++ b/client/ssh/client/terminal_windows.go @@ -10,6 +10,8 @@ import ( log "github.com/sirupsen/logrus" "golang.org/x/crypto/ssh" + + nbssh "github.com/netbirdio/netbird/client/ssh" ) const ( @@ -80,28 +82,14 @@ func (c *Client) setupTerminalMode(_ context.Context, session *ssh.Session) erro w, h := c.getWindowsConsoleSize() modes := ssh.TerminalModes{ - ssh.ECHO: 1, - ssh.TTY_OP_ISPEED: 14400, - ssh.TTY_OP_OSPEED: 14400, - ssh.ICRNL: 1, - ssh.OPOST: 1, - ssh.ONLCR: 1, - ssh.ISIG: 1, - ssh.ICANON: 1, - ssh.VINTR: 3, // Ctrl+C - ssh.VQUIT: 28, // Ctrl+\ - ssh.VERASE: 127, // Backspace - ssh.VKILL: 21, // Ctrl+U - ssh.VEOF: 4, // Ctrl+D - ssh.VEOL: 0, - ssh.VEOL2: 0, - ssh.VSTART: 17, // Ctrl+Q - ssh.VSTOP: 19, // Ctrl+S - ssh.VSUSP: 26, // Ctrl+Z - ssh.VDISCARD: 15, // Ctrl+O - ssh.VWERASE: 23, // Ctrl+W - ssh.VLNEXT: 22, // Ctrl+V - ssh.VREPRINT: 18, // Ctrl+R + ssh.ICRNL: 1, + ssh.OPOST: 1, + ssh.ONLCR: 1, + ssh.ISIG: 1, + ssh.ICANON: 1, + } + for mode, value := range nbssh.DefaultTerminalModes { + modes[mode] = value } if err := session.RequestPty("xterm-256color", h, w, modes); err != nil { diff --git a/client/ssh/common.go b/client/ssh/common.go index 92e647b7d..4ebf8842a 100644 --- a/client/ssh/common.go +++ b/client/ssh/common.go @@ -13,6 +13,7 @@ import ( "golang.org/x/crypto/ssh" "github.com/netbirdio/netbird/client/proto" + "github.com/netbirdio/netbird/util" ) const ( @@ -34,6 +35,19 @@ type HostKeyVerifier interface { VerifySSHHostKey(peerAddress string, key []byte) error } +// PeerKeyLookup returns the stored SSH host key for a peer address. +type PeerKeyLookup func(peerAddress string) ([]byte, bool) + +// VerifySSHHostKey implements HostKeyVerifier by looking up the stored key +// and comparing it against the presented key. +func (l PeerKeyLookup) VerifySSHHostKey(peerAddress string, presentedKey []byte) error { + storedKey, found := l(peerAddress) + if !found { + return ErrPeerNotFound + } + return VerifyHostKey(storedKey, presentedKey, peerAddress) +} + // DaemonHostKeyVerifier implements HostKeyVerifier using the NetBird daemon type DaemonHostKeyVerifier struct { client proto.DaemonServiceClient @@ -92,7 +106,8 @@ func printAuthInstructions(stderr io.Writer, authResponse *proto.RequestJWTAuthR // RequestJWTToken requests or retrieves a JWT token for SSH authentication func RequestJWTToken(ctx context.Context, client proto.DaemonServiceClient, stdout, stderr io.Writer, useCache bool, hint string, openBrowser func(string) error) (string, error) { - req := &proto.RequestJWTAuthRequest{} + // the ssh client runs in the user's session, the daemon does not: tell it what we can see + req := &proto.RequestJWTAuthRequest{HasGraphicalSession: util.HasGraphicalSession()} if hint != "" { req.Hint = &hint } @@ -193,4 +208,3 @@ func buildAddressList(hostname string, remote net.Addr) []string { } return addresses } - diff --git a/client/ssh/handshake.go b/client/ssh/handshake.go new file mode 100644 index 000000000..e78a806be --- /dev/null +++ b/client/ssh/handshake.go @@ -0,0 +1,45 @@ +package ssh + +import ( + "context" + "fmt" + "io" + "net" + "time" + + log "github.com/sirupsen/logrus" + "golang.org/x/crypto/ssh" +) + +// Handshake runs the SSH client handshake on an already dialed conn and +// returns the resulting client. Dialing bounds only the TCP establishment; +// without a deadline on the socket a peer that accepts and then goes silent +// blocks the handshake forever, so the context deadline is applied to conn +// for the duration of the handshake. conn is closed on any error. +func Handshake(ctx context.Context, conn net.Conn, addr string, config *ssh.ClientConfig) (*ssh.Client, error) { + if deadline, ok := ctx.Deadline(); ok { + if err := conn.SetDeadline(deadline); err != nil { + closeHandshake(conn, "conn after deadline error") + return nil, fmt.Errorf("set handshake deadline: %w", err) + } + } + + sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config) + if err != nil { + closeHandshake(conn, "conn after handshake error") + return nil, fmt.Errorf("ssh handshake: %w", err) + } + + if err := conn.SetDeadline(time.Time{}); err != nil { + closeHandshake(sshConn, "ssh conn after deadline clear error") + return nil, fmt.Errorf("clear handshake deadline: %w", err) + } + + return ssh.NewClient(sshConn, chans, reqs), nil +} + +func closeHandshake(c io.Closer, label string) { + if err := c.Close(); err != nil { + log.Debugf("ssh: close %s: %v", label, err) + } +} diff --git a/client/ssh/proxy/proxy.go b/client/ssh/proxy/proxy.go index 73b50122c..070515b57 100644 --- a/client/ssh/proxy/proxy.go +++ b/client/ssh/proxy/proxy.go @@ -9,7 +9,6 @@ import ( "net" "os" "strconv" - "strings" "sync" "time" @@ -17,8 +16,8 @@ import ( log "github.com/sirupsen/logrus" cryptossh "golang.org/x/crypto/ssh" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/proto" nbssh "github.com/netbirdio/netbird/client/ssh" @@ -55,8 +54,8 @@ type SSHProxy struct { } func New(daemonAddr, targetHost string, targetPort int, stderr io.Writer, browserOpener func(string) error) (*SSHProxy, error) { - grpcAddr := strings.TrimPrefix(daemonAddr, "tcp://") - grpcConn, err := grpc.NewClient(grpcAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) + target, opts := daemonaddr.DialTarget(daemonAddr) + grpcConn, err := grpc.NewClient(target, opts...) if err != nil { return nil, fmt.Errorf("connect to daemon: %w", err) } @@ -611,13 +610,10 @@ func (p *SSHProxy) dialBackend(ctx context.Context, addr, user, jwtToken string) return nil, fmt.Errorf("connect to server: %w", err) } - clientConn, chans, reqs, err := cryptossh.NewClientConn(conn, addr, config) - if err != nil { - _ = conn.Close() - return nil, fmt.Errorf("SSH handshake: %w", err) - } + handshakeCtx, cancel := context.WithTimeout(ctx, sshHandshakeTimeout) + defer cancel() - return cryptossh.NewClient(clientConn, chans, reqs), nil + return nbssh.Handshake(handshakeCtx, conn, addr, config) } func (p *SSHProxy) verifyHostKey(hostname string, remote net.Addr, key cryptossh.PublicKey) error { diff --git a/client/ssh/server/command_execution.go b/client/ssh/server/command_execution.go index b0a85fe4b..c8b3240d0 100644 --- a/client/ssh/server/command_execution.go +++ b/client/ssh/server/command_execution.go @@ -75,8 +75,8 @@ func (s *Server) createCommand(logger *log.Entry, privilegeResult PrivilegeCheck } // Try su first for system integration (PAM/audit) when privileged - cmd, err := s.createSuCommand(logger, session, localUser, hasPty) - if err != nil || privilegeResult.UsedFallback { + cmd, err := s.createSuCommand(logger, session, localUser, hasPty) //nolint:staticcheck + if err != nil || privilegeResult.UsedFallback { //nolint:staticcheck // always errors on platforms without su logger.Debugf("su command failed, falling back to executor: %v", err) cmd, cleanup, err := s.createExecutorCommand(logger, session, localUser, hasPty) if err != nil { diff --git a/client/ssh/server/command_execution_windows.go b/client/ssh/server/command_execution_windows.go index e1ba777f6..feb8daa26 100644 --- a/client/ssh/server/command_execution_windows.go +++ b/client/ssh/server/command_execution_windows.go @@ -243,7 +243,7 @@ func (s *Server) setUserEnvironmentVariables(envMap map[string]string, userProfi // prepareCommandEnv prepares environment variables for command execution on Windows func (s *Server) prepareCommandEnv(logger *log.Entry, localUser *user.User, session ssh.Session) []string { - username, domain := s.parseUsername(localUser.Username) + username, domain := parseUsername(localUser.Username) userEnv, err := s.getUserEnvironment(logger, username, domain) if err != nil { log.Debugf("failed to get user environment for %s\\%s, using fallback: %v", domain, username, err) @@ -383,7 +383,7 @@ func (s *Server) executeCommandWithPty(logger *log.Entry, session ssh.Session, _ return false } - username, domain := s.parseUsername(localUser.Username) + username, domain := parseUsername(localUser.Username) shell := getUserShell(localUser.Uid) req := PtyExecutionRequest{ diff --git a/client/ssh/server/port_forwarding.go b/client/ssh/server/port_forwarding.go index a47fdb48a..81d3b9173 100644 --- a/client/ssh/server/port_forwarding.go +++ b/client/ssh/server/port_forwarding.go @@ -133,7 +133,12 @@ func (s *Server) checkPrivilegedPortAccess(forwardType string, port uint32, resu return nil } - if result.User != nil && isPrivilegedUsername(result.User.Username) { + // Only uid 0 may bind below the threshold, which is the kernel's own rule and + // is asked directly rather than through isPrivilegedOrUnknown: that helper + // reports an account it cannot evaluate as privileged, which is safe for a + // refusal and unsafe for a grant such as this one. Windows has returned + // above, so Uid here is a Unix uid and never a SID. + if result.User != nil && result.User.Uid == "0" { return nil } diff --git a/client/ssh/server/privileges_other.go b/client/ssh/server/privileges_other.go new file mode 100644 index 000000000..89440dea8 --- /dev/null +++ b/client/ssh/server/privileges_other.go @@ -0,0 +1,16 @@ +//go:build !windows + +package server + +// isProcessElevated is only meaningful on Windows; other platforms use the +// effective UID check in isCurrentProcessPrivileged. +func isProcessElevated() bool { + return false +} + +// isWindowsAccountPrivilegedOrUnknown is only reachable on Windows. Report +// privileged on other platforms so a caller refusing privileged accounts fails +// closed. +func isWindowsAccountPrivilegedOrUnknown(string) bool { + return true +} diff --git a/client/ssh/server/privileges_windows.go b/client/ssh/server/privileges_windows.go new file mode 100644 index 000000000..41ea00fd7 --- /dev/null +++ b/client/ssh/server/privileges_windows.go @@ -0,0 +1,228 @@ +//go:build windows + +package server + +import ( + "fmt" + "strings" + "unsafe" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" +) + +var ( + netapi32 = windows.NewLazySystemDLL("netapi32.dll") + procNetUserGetLocalGroups = netapi32.NewProc("NetUserGetLocalGroups") +) + +const ( + // lgIncludeIndirect makes NetUserGetLocalGroups also return local groups + // the user belongs to through a global group. + lgIncludeIndirect = 0x1 + maxPreferredLength = 0xFFFFFFFF +) + +// localGroupUsersInfo0 mirrors LOCALGROUP_USERS_INFO_0. +type localGroupUsersInfo0 struct { + name *uint16 +} + +// isProcessElevated reports whether the current process token is elevated +// (TokenElevation): true for elevated administrators, the built-in +// Administrator, administrators with UAC disabled, and SYSTEM; false for +// standard users and administrators running with a UAC-filtered token. +func isProcessElevated() bool { + return windows.GetCurrentProcessToken().IsElevated() +} + +// isWindowsAccountPrivilegedOrUnknown reports whether the account is privileged +// on this machine: a well-known service account, a built-in Administrator +// (RID 500), or a member of the local Administrators group, directly or through +// nested groups. +// +// An account whose privilege cannot be determined counts as privileged, which +// is why the name says "or unknown". That is fail-closed for a caller that +// refuses privileged accounts, and fail-open for a caller that grants something +// to them, so only the former may use this. +func isWindowsAccountPrivilegedOrUnknown(username string) bool { + sid, _, _, err := windows.LookupSID("", username) + if err != nil { + log.Warnf("privilege check: SID lookup for %q failed, treating as privileged: %v", username, err) + return true + } + + if isPrivilegedUserSID(sid) { + return true + } + + member, err := isLocalAdminsMember(username) + if err != nil { + log.Warnf("privilege check: cannot determine Administrators membership for %q, treating as privileged: %v", username, err) + return true + } + return member +} + +// isPrivilegedUserSID reports whether the SID itself identifies a privileged +// principal, without consulting group membership. +func isPrivilegedUserSID(sid *windows.SID) bool { + wellKnown := []windows.WELL_KNOWN_SID_TYPE{ + windows.WinLocalSystemSid, + windows.WinLocalServiceSid, + windows.WinNetworkServiceSid, + windows.WinBuiltinAdministratorsSid, + } + for _, sidType := range wellKnown { + if sid.IsWellKnown(sidType) { + return true + } + } + return isBuiltinAdministratorSID(sid) +} + +// isBuiltinAdministratorSID reports whether the SID is a machine or domain +// built-in Administrator account (S-1-5-21-...-500). RID 500 is reserved for +// that account; it can be renamed but cannot be removed from the +// Administrators group. +func isBuiltinAdministratorSID(sid *windows.SID) bool { + if sid.IdentifierAuthority() != windows.SECURITY_NT_AUTHORITY { + return false + } + count := sid.SubAuthorityCount() + if count < 2 || sid.SubAuthority(0) != 21 { + return false + } + return sid.SubAuthority(uint32(count-1)) == 500 +} + +// isLocalAdminsMember reports whether the account is a member of the local +// Administrators group. +// +// Local accounts are checked against the local SAM, which is authoritative for +// them and, unlike a token, cannot under-report: UAC filters the tokens of +// local administrators, and a filtered token carries Administrators as +// deny-only, which a membership check on the token would read as "not a +// member". Domain accounts are exempt from that filtering, so for them an S4U +// token is preferred because its group list is LSA's transitive expansion and +// therefore covers nested and universal groups plus the machine's own local +// groups. NetUserGetLocalGroups expands only one global-group hop but needs no +// logon, so it serves as the fallback when no token can be obtained. +func isLocalAdminsMember(username string) (bool, error) { + adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid) + if err != nil { + return false, fmt.Errorf("create Administrators SID: %w", err) + } + + account, domain := parseUsername(username) + if NewPrivilegeDropper().isLocalUser(domain) { + return localGroupsContainSID(account, adminSid) + } + + member, s4uErr := s4uTokenIsMember(account, domain, adminSid) + if s4uErr == nil { + return member, nil + } + log.Debugf("privilege check: S4U membership check for %q failed, falling back to local group enumeration: %v", username, s4uErr) + + member, err = localGroupsContainSID(buildUserCpn(account, domain), adminSid) + if err != nil { + return false, fmt.Errorf("S4U check: %w; local group enumeration: %w", s4uErr, err) + } + return member, nil +} + +// s4uTokenIsMember obtains an S4U token for the account and checks whether the +// given SID is enabled in it. +func s4uTokenIsMember(account, domain string, sid *windows.SID) (bool, error) { + token, err := generateS4UUserToken(log.NewEntry(log.StandardLogger()), account, domain) + if err != nil { + return false, err + } + defer func() { + if err := windows.CloseHandle(token); err != nil { + log.Debugf("close S4U token: %v", err) + } + }() + return windows.Token(token).IsMember(sid) +} + +// localGroupsContainSID reports whether the wanted group is among the local +// groups the account belongs to, directly or through a global group. +// +// The wanted SID is resolved to its group name once and compared against the +// enumerated names. Well-known SIDs resolve from a static table, so that lookup +// needs no domain controller, and it keeps the comparison correct for a renamed +// or localized group because both sides then carry the new name. Resolving each +// enumerated name back to a SID instead would add a lookup per group that can +// block until it times out while a domain controller is unreachable, and cannot +// change the outcome: the names enumerated here are local groups of this +// machine, whose names are unique, so a name match identifies the group. +// +// A failure to resolve the wanted SID is returned rather than reported as +// "not a member", so a privilege check built on this fails closed. +func localGroupsContainSID(username string, want *windows.SID) (bool, error) { + wantName, _, _, err := want.LookupAccount("") + if err != nil { + return false, fmt.Errorf("resolve group SID %s to a name: %w", want, err) + } + + groups, err := netUserGetLocalGroups(username) + if err != nil { + return false, err + } + + for _, group := range groups { + if strings.EqualFold(group, wantName) { + return true, nil + } + } + return false, nil +} + +// netUserGetLocalGroups returns the names of the local groups the account is a +// member of, including indirect membership through global groups. +func netUserGetLocalGroups(username string) ([]string, error) { + name16, err := windows.UTF16PtrFromString(username) + if err != nil { + return nil, fmt.Errorf("convert username: %w", err) + } + + var buf *byte + var entriesRead, totalEntries uint32 + status, _, _ := procNetUserGetLocalGroups.Call( + 0, // local server + uintptr(unsafe.Pointer(name16)), + 0, // level 0: LOCALGROUP_USERS_INFO_0 + lgIncludeIndirect, + uintptr(unsafe.Pointer(&buf)), + maxPreferredLength, + uintptr(unsafe.Pointer(&entriesRead)), + uintptr(unsafe.Pointer(&totalEntries)), + ) + if status != 0 { + return nil, fmt.Errorf("NetUserGetLocalGroups for %q: status %d", username, status) + } + if buf == nil { + return nil, nil + } + defer func() { + if err := windows.NetApiBufferFree(buf); err != nil { + log.Debugf("free NetApi buffer: %v", err) + } + }() + + // MAX_PREFERRED_LENGTH makes the API allocate as much as it needs, so a + // short read is not expected. Report it rather than silently returning a + // subset of the account's groups. + if entriesRead != totalEntries { + return nil, fmt.Errorf("NetUserGetLocalGroups for %q returned %d of %d groups", username, entriesRead, totalEntries) + } + + entries := unsafe.Slice((*localGroupUsersInfo0)(unsafe.Pointer(buf)), entriesRead) + groups := make([]string, 0, entriesRead) + for _, entry := range entries { + groups = append(groups, windows.UTF16PtrToString(entry.name)) + } + return groups, nil +} diff --git a/client/ssh/server/privileges_windows_test.go b/client/ssh/server/privileges_windows_test.go new file mode 100644 index 000000000..983fccdf7 --- /dev/null +++ b/client/ssh/server/privileges_windows_test.go @@ -0,0 +1,293 @@ +//go:build windows + +package server + +import ( + "os/user" + "testing" + "unsafe" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/sys/windows" +) + +// filterNormalAccount limits NetUserEnum to normal user accounts. +const filterNormalAccount = 0x2 + +// TOKEN_ELEVATION_TYPE values. +const ( + tokenElevationTypeDefault = 1 + tokenElevationTypeFull = 2 + tokenElevationTypeLimited = 3 +) + +// tokenElevationType reads TokenElevationType from a token. +func tokenElevationType(token windows.Token) (uint32, error) { + var elevationType, returnedLen uint32 + err := windows.GetTokenInformation(token, windows.TokenElevationType, + (*byte)(unsafe.Pointer(&elevationType)), uint32(unsafe.Sizeof(elevationType)), &returnedLen) + if err != nil { + return 0, err + } + return elevationType, nil +} + +// userInfo0 mirrors USER_INFO_0. +type userInfo0 struct { + name *uint16 +} + +func mustParseSID(t *testing.T, s string) *windows.SID { + t.Helper() + sid, err := windows.StringToSid(s) + require.NoError(t, err, "parse SID %s", s) + return sid +} + +// localAccountNames returns the names of the local user accounts. +func localAccountNames(t *testing.T) []string { + t.Helper() + + var buf *byte + var entriesRead, totalEntries, resume uint32 + err := windows.NetUserEnum(nil, 0, filterNormalAccount, &buf, maxPreferredLength, + &entriesRead, &totalEntries, &resume) + require.NoError(t, err, "enumerate local users") + t.Cleanup(func() { + require.NoError(t, windows.NetApiBufferFree(buf), "free NetApi buffer") + }) + + entries := unsafe.Slice((*userInfo0)(unsafe.Pointer(buf)), entriesRead) + names := make([]string, 0, entriesRead) + for _, entry := range entries { + names = append(names, windows.UTF16PtrToString(entry.name)) + } + return names +} + +// localAccountNameByRID returns the name of the local account carrying the +// given RID. Accounts such as Administrator and Guest can be renamed and are +// localized, so tests must not name them literally. +func localAccountNameByRID(t *testing.T, rid uint32) string { + t.Helper() + + for _, name := range localAccountNames(t) { + sid, _, _, err := windows.LookupSID("", name) + if err != nil { + continue + } + if sid.IdentifierAuthority() != windows.SECURITY_NT_AUTHORITY { + continue + } + count := sid.SubAuthorityCount() + if count < 2 || sid.SubAuthority(0) != 21 { + continue + } + if sid.SubAuthority(uint32(count-1)) == rid { + return name + } + } + + t.Fatalf("no local account with RID %d", rid) + return "" +} + +// wellKnownAccountName resolves a well-known SID to the qualified account name +// the local system uses for it, which is localized. +func wellKnownAccountName(t *testing.T, sidType windows.WELL_KNOWN_SID_TYPE) string { + t.Helper() + + sid, err := windows.CreateWellKnownSid(sidType) + require.NoError(t, err, "create well-known SID") + name, domain, _, err := sid.LookupAccount("") + require.NoError(t, err, "resolve %s to an account name", sid) + if domain == "" { + return name + } + return domain + `\` + name +} + +func TestIsBuiltinAdministratorSID(t *testing.T) { + tests := []struct { + name string + sid string + want bool + }{ + {"machine_administrator", "S-1-5-21-1111111111-2222222222-3333333333-500", true}, + {"domain_administrator", "S-1-5-21-3390233681-4087452608-412898826-500", true}, + {"regular_user", "S-1-5-21-1111111111-2222222222-3333333333-1001", false}, + {"guest_account", "S-1-5-21-1111111111-2222222222-3333333333-501", false}, + {"domain_admins_group", "S-1-5-21-1111111111-2222222222-3333333333-512", false}, + {"system", "S-1-5-18", false}, + {"administrators_group", "S-1-5-32-544", false}, + {"non_nt_authority", "S-1-1-0", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := isBuiltinAdministratorSID(mustParseSID(t, tt.sid)) + assert.Equal(t, tt.want, result, "RID 500 detection for %s", tt.sid) + }) + } +} + +func TestIsPrivilegedUserSID(t *testing.T) { + tests := []struct { + name string + sid string + want bool + }{ + {"local_system", "S-1-5-18", true}, + {"local_service", "S-1-5-19", true}, + {"network_service", "S-1-5-20", true}, + {"administrators_group", "S-1-5-32-544", true}, + {"builtin_administrator", "S-1-5-21-1111111111-2222222222-3333333333-500", true}, + {"regular_user", "S-1-5-21-1111111111-2222222222-3333333333-1001", false}, + {"users_group", "S-1-5-32-545", false}, + {"everyone", "S-1-1-0", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := isPrivilegedUserSID(mustParseSID(t, tt.sid)) + assert.Equal(t, tt.want, result, "SID privilege classification for %s", tt.sid) + }) + } +} + +func TestIsWindowsAccountPrivilegedOrUnknown(t *testing.T) { + tests := []struct { + name string + username string + want bool + }{ + {"system", wellKnownAccountName(t, windows.WinLocalSystemSid), true}, + {"local_service", wellKnownAccountName(t, windows.WinLocalServiceSid), true}, + {"network_service", wellKnownAccountName(t, windows.WinNetworkServiceSid), true}, + {"administrators_group", wellKnownAccountName(t, windows.WinBuiltinAdministratorsSid), true}, + // The built-in Administrator (RID 500) and Guest (RID 501) accounts + // exist on every Windows installation, though they may be disabled. + {"builtin_administrator", localAccountNameByRID(t, 500), true}, + {"guest", localAccountNameByRID(t, 501), false}, + // Unresolvable accounts fail closed. + {"nonexistent_user", "netbird-no-such-user", true}, + {"empty_username", "", true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := isWindowsAccountPrivilegedOrUnknown(tt.username) + assert.Equal(t, tt.want, result, "account privilege classification for %q", tt.username) + }) + } +} + +func TestIsProcessElevated(t *testing.T) { + elevated := isProcessElevated() + + // TokenElevationType is a second, independent view of the same token: + // Full means elevated and Limited means a filtered administrator, while + // Default covers both a standard user and an administrator with no linked + // token (UAC off, the built-in Administrator, SYSTEM), so it implies nothing. + elevationType, err := tokenElevationType(windows.GetCurrentProcessToken()) + require.NoError(t, err, "read token elevation type") + + adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid) + require.NoError(t, err, "create Administrators SID") + + // Token(0) makes CheckTokenMembership evaluate the caller's own token. It + // counts only enabled SIDs, so a filtered administrator reports false here. + member, err := windows.Token(0).IsMember(adminSid) + require.NoError(t, err, "check own Administrators membership") + + t.Logf("elevated=%v elevationType=%d memberOfAdministrators=%v", elevated, elevationType, member) + + switch elevationType { + case tokenElevationTypeFull: + assert.True(t, elevated, "a token of elevation type Full must report elevated") + case tokenElevationTypeLimited: + assert.False(t, elevated, "a filtered administrator token must not report elevated") + } + + // Administrators enabled in the token means the token wields administrative + // rights, which is what elevation reports. + if member { + assert.True(t, elevated, "token with enabled Administrators membership must report elevated") + } +} + +// TestS4UMembershipAgreesWithLocalGroups exercises the S4U token path used +// for domain accounts. S4U logons need the TCB privilege, so the test runs +// only as SYSTEM (which is how CI executes the suite). For local accounts the +// token's Administrators membership must agree with the SAM enumeration. +func TestS4UMembershipAgreesWithLocalGroups(t *testing.T) { + system, err := windows.CreateWellKnownSid(windows.WinLocalSystemSid) + require.NoError(t, err, "create SYSTEM SID") + current, err := user.Current() + require.NoError(t, err, "get current user") + if current.Uid != system.String() { + t.Skipf("S4U logon requires SYSTEM (running as %s)", current.Username) + } + + adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid) + require.NoError(t, err, "create Administrators SID") + + checked := 0 + for _, name := range localAccountNames(t) { + viaToken, err := s4uTokenIsMember(name, ".", adminSid) + if err != nil { + // Disabled or logon-restricted accounts cannot get an S4U logon. + t.Logf("skipping %s: %v", name, err) + continue + } + viaSAM, err := localGroupsContainSID(name, adminSid) + require.NoError(t, err, "enumerate local groups for %s", name) + + assert.Equal(t, viaSAM, viaToken, "S4U token and SAM enumeration must agree on Administrators membership for %s", name) + checked++ + } + // Ineligible accounts are skipped, so without this the test could report + // success while comparing nothing at all. + require.Positive(t, checked, "no local account completed an S4U logon, so nothing was compared") + t.Logf("checked %d local accounts via S4U", checked) +} + +// TestLocalGroupsContainSID_Administrator checks the positive case against the +// built-in Administrator, a member of Administrators on every installation. +func TestLocalGroupsContainSID_Administrator(t *testing.T) { + adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid) + require.NoError(t, err, "create Administrators SID") + + administrator := localAccountNameByRID(t, 500) + member, err := localGroupsContainSID(administrator, adminSid) + require.NoError(t, err, "enumerate local groups for %s", administrator) + assert.True(t, member, "%s is a member of the Administrators group", administrator) +} + +// TestLocalGroupsContainSID_UnresolvableGroupFailsClosed covers a wanted SID +// that resolves to no group: the error must surface rather than being reported +// as "not a member", so the privilege check treats the account as privileged. +func TestLocalGroupsContainSID_UnresolvableGroupFailsClosed(t *testing.T) { + unknown := mustParseSID(t, "S-1-5-21-1111111111-2222222222-3333333333-4444") + + _, err := localGroupsContainSID(localAccountNameByRID(t, 500), unknown) + require.Error(t, err, "must report an error when the wanted group cannot be identified") +} + +func TestLocalGroupsContainSID_Guest(t *testing.T) { + guestsSid, err := windows.CreateWellKnownSid(windows.WinBuiltinGuestsSid) + require.NoError(t, err, "create Guests SID") + adminsSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid) + require.NoError(t, err, "create Administrators SID") + + guest := localAccountNameByRID(t, 501) + + inGuests, err := localGroupsContainSID(guest, guestsSid) + require.NoError(t, err, "enumerate local groups for %s", guest) + assert.True(t, inGuests, "%s is a member of the Guests group", guest) + + inAdmins, err := localGroupsContainSID(guest, adminsSid) + require.NoError(t, err, "enumerate local groups for %s", guest) + assert.False(t, inAdmins, "%s is not a member of the Administrators group", guest) +} diff --git a/client/ssh/server/server_config_test.go b/client/ssh/server/server_config_test.go index f70e29963..983bd7a43 100644 --- a/client/ssh/server/server_config_test.go +++ b/client/ssh/server/server_config_test.go @@ -239,6 +239,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { forwardType string port uint32 username string + uid string expectError bool errorMsg string skipOnWindows bool @@ -248,6 +249,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { forwardType: "remote", port: 80, username: "testuser", + uid: "1000", expectError: true, errorMsg: "cannot bind to privileged port", skipOnWindows: true, @@ -257,6 +259,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { forwardType: "tcpip-forward", port: 443, username: "testuser", + uid: "1000", expectError: true, errorMsg: "cannot bind to privileged port", skipOnWindows: true, @@ -266,6 +269,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { forwardType: "remote", port: 8080, username: "testuser", + uid: "1000", expectError: false, }, { @@ -273,6 +277,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { forwardType: "remote", port: 0, username: "testuser", + uid: "1000", expectError: false, }, { @@ -280,13 +285,35 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { forwardType: "remote", port: 22, username: "root", + uid: "0", expectError: false, }, + { + // Only uid 0 is privileged, whatever the account is called. + name: "uid 0 under another name may bind a privileged port", + forwardType: "remote", + port: 22, + username: "toor", + uid: "0", + expectError: false, + skipOnWindows: true, + }, + { + name: "account named root without uid 0 may not", + forwardType: "remote", + port: 22, + username: "root", + uid: "1000", + expectError: true, + errorMsg: "cannot bind to privileged port", + skipOnWindows: true, + }, { name: "local forward privileged port allowed for non-root", forwardType: "local", port: 80, username: "testuser", + uid: "1000", expectError: false, }, } @@ -299,7 +326,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) { result := PrivilegeCheckResult{ Allowed: true, - User: &user.User{Username: tt.username}, + User: &user.User{Username: tt.username, Uid: tt.uid}, } err := server.checkPrivilegedPortAccess(tt.forwardType, tt.port, result) @@ -420,6 +447,13 @@ func TestServer_PortConflictHandling(t *testing.T) { func TestServer_IsPrivilegedUser(t *testing.T) { + // Windows classification depends on account SIDs and group membership, and + // the accounts involved carry localized, renameable names. It is covered by + // TestIsWindowsAccountPrivileged, which resolves them from well-known SIDs. + if runtime.GOOS == "windows" { + t.Skip("covered by TestIsWindowsAccountPrivileged") + } + tests := []struct { username string expected bool @@ -440,44 +474,16 @@ func TestServer_IsPrivilegedUser(t *testing.T) { expected: false, description: "empty username should not be privileged", }, - } - - // Add Windows-specific tests - if runtime.GOOS == "windows" { - tests = append(tests, []struct { - username string - expected bool - description string - }{ - { - username: "Administrator", - expected: true, - description: "Administrator should be considered privileged on Windows", - }, - { - username: "administrator", - expected: true, - description: "administrator should be considered privileged on Windows (case insensitive)", - }, - }...) - } else { - // On non-Windows systems, Administrator should not be privileged - tests = append(tests, []struct { - username string - expected bool - description string - }{ - { - username: "Administrator", - expected: false, - description: "Administrator should not be privileged on non-Windows systems", - }, - }...) + { + username: "Administrator", + expected: false, + description: "Administrator should not be privileged on non-Windows systems", + }, } for _, tt := range tests { t.Run(tt.description, func(t *testing.T) { - result := isPrivilegedUsername(tt.username) + result := isPrivilegedOrUnknown(tt.username) assert.Equal(t, tt.expected, result, tt.description) }) } diff --git a/client/ssh/server/sftp_windows.go b/client/ssh/server/sftp_windows.go index dc532b9e7..25cd17298 100644 --- a/client/ssh/server/sftp_windows.go +++ b/client/ssh/server/sftp_windows.go @@ -17,7 +17,7 @@ import ( // createSftpCommand creates a Windows SFTP command with user switching. // The caller must close the returned token handle after starting the process. func (s *Server) createSftpCommand(targetUser *user.User, sess ssh.Session) (*exec.Cmd, windows.Token, error) { - username, domain := s.parseUsername(targetUser.Username) + username, domain := parseUsername(targetUser.Username) netbirdPath, err := os.Executable() if err != nil { diff --git a/client/ssh/server/user_utils.go b/client/ssh/server/user_utils.go index bc2aa2d7d..6c8142b30 100644 --- a/client/ssh/server/user_utils.go +++ b/client/ssh/server/user_utils.go @@ -16,11 +16,6 @@ var ( ErrPrivilegedUserSwitch = errors.New("cannot switch to privileged user - current user lacks required privileges") ) -// isPlatformUnix returns true for Unix-like platforms (Linux, macOS, etc.) -func isPlatformUnix() bool { - return getCurrentOS() != "windows" -} - // Dependency injection variables for testing - allows mocking dynamic runtime checks var ( getCurrentUser = currentUserWithGetent @@ -29,6 +24,9 @@ var ( getIsProcessPrivileged = isCurrentProcessPrivileged getEuid = os.Geteuid + + getProcessElevated = isProcessElevated + getWindowsAccountPrivilegedOrUnknown = isWindowsAccountPrivilegedOrUnknown ) const ( @@ -65,6 +63,13 @@ type PrivilegeCheckResult struct { RequiresUserSwitching bool } +// privilegeCheckContext holds all context needed for privilege checking +type privilegeCheckContext struct { + currentUser *user.User + currentUserPrivileged bool + allowRoot bool +} + // CheckPrivileges performs comprehensive privilege checking for all SSH features. // This is the single source of truth for privilege decisions across the SSH server. func (s *Server) CheckPrivileges(req PrivilegeCheckRequest) PrivilegeCheckResult { @@ -75,7 +80,7 @@ func (s *Server) CheckPrivileges(req PrivilegeCheckRequest) PrivilegeCheckResult // Handle empty username case - but still check root access controls if req.RequestedUsername == "" { - if isPrivilegedUsername(context.currentUser.Username) && !context.allowRoot { + if isPrivilegedOrUnknown(context.currentUser.Username) && !context.allowRoot { return PrivilegeCheckResult{ Allowed: false, Error: &PrivilegedUserError{Username: context.currentUser.Username}, @@ -135,7 +140,7 @@ func (s *Server) checkUserRequest(ctx *privilegeCheckContext, req PrivilegeCheck needsUserSwitching := !isSameResolvedUser(resolvedUser, ctx.currentUser) - if isPrivilegedUsername(resolvedUser.Username) && !ctx.allowRoot { + if isPrivilegedOrUnknown(resolvedUser.Username) && !ctx.allowRoot { return PrivilegeCheckResult{ Allowed: false, Error: &PrivilegedUserError{Username: resolvedUser.Username}, @@ -175,6 +180,42 @@ func (s *Server) resolveRequestedUser(requestedUsername string) (*user.User, err return u, nil } +// SetAllowRootLogin configures root login access +func (s *Server) SetAllowRootLogin(allow bool) { + s.mu.Lock() + defer s.mu.Unlock() + s.allowRootLogin = allow +} + +// userNameLookup performs user lookup with root login permission check +func (s *Server) userNameLookup(username string) (*user.User, error) { + result, err := s.userPrivilegeCheck(username) + if err != nil { + return nil, err + } + return result.User, nil +} + +// userPrivilegeCheck performs user lookup with full privilege check result +func (s *Server) userPrivilegeCheck(username string) (PrivilegeCheckResult, error) { + result := s.CheckPrivileges(PrivilegeCheckRequest{ + RequestedUsername: username, + FeatureSupportsUserSwitch: true, + FeatureName: FeatureSSHLogin, + }) + + if !result.Allowed { + return result, result.Error + } + + return result, nil +} + +// isPlatformUnix returns true for Unix-like platforms (Linux, macOS, etc.) +func isPlatformUnix() bool { + return getCurrentOS() != "windows" +} + // isSameResolvedUser compares two resolved user identities func isSameResolvedUser(user1, user2 *user.User) bool { if user1 == nil || user2 == nil { @@ -183,13 +224,6 @@ func isSameResolvedUser(user1, user2 *user.User) bool { return user1.Uid == user2.Uid } -// privilegeCheckContext holds all context needed for privilege checking -type privilegeCheckContext struct { - currentUser *user.User - currentUserPrivileged bool - allowRoot bool -} - // isSameUser checks if two usernames refer to the same user // SECURITY: This function must be conservative - it should only return true // when we're certain both usernames refer to the exact same user identity @@ -253,159 +287,30 @@ func isWindowsSameUser(requestedUsername, currentUsername string) bool { return strings.EqualFold(reqDomain, curDomain) } -// SetAllowRootLogin configures root login access -func (s *Server) SetAllowRootLogin(allow bool) { - s.mu.Lock() - defer s.mu.Unlock() - s.allowRootLogin = allow -} - -// userNameLookup performs user lookup with root login permission check -func (s *Server) userNameLookup(username string) (*user.User, error) { - result := s.CheckPrivileges(PrivilegeCheckRequest{ - RequestedUsername: username, - FeatureSupportsUserSwitch: true, - FeatureName: FeatureSSHLogin, - }) - - if !result.Allowed { - return nil, result.Error - } - - return result.User, nil -} - -// userPrivilegeCheck performs user lookup with full privilege check result -func (s *Server) userPrivilegeCheck(username string) (PrivilegeCheckResult, error) { - result := s.CheckPrivileges(PrivilegeCheckRequest{ - RequestedUsername: username, - FeatureSupportsUserSwitch: true, - FeatureName: FeatureSSHLogin, - }) - - if !result.Allowed { - return result, result.Error - } - - return result, nil -} - -// isPrivilegedUsername checks if the given username represents a privileged user across platforms. -// On Unix: root -// On Windows: Administrator, SYSTEM (case-insensitive) -// Handles domain-qualified usernames like "DOMAIN\Administrator" or "user@domain.com" -func isPrivilegedUsername(username string) bool { +// isPrivilegedOrUnknown reports whether the given username represents a +// privileged user, or on Windows an account whose privilege could not be +// determined. +// On Unix: root. +// On Windows: well-known service accounts, built-in Administrator accounts, +// and members of the local Administrators group; handles domain-qualified +// usernames like "DOMAIN\user" or "user@domain.com". An account that cannot be +// resolved or evaluated is reported as privileged. +// +// Use this to refuse privileged accounts, never to grant them anything: the +// undetermined case is safe for a refusal and unsafe for a grant. +func isPrivilegedOrUnknown(username string) bool { if getCurrentOS() != "windows" { return username == "root" } - - bareUsername := username - // Handle Windows domain format: DOMAIN\username - if idx := strings.LastIndex(username, `\`); idx != -1 { - bareUsername = username[idx+1:] - } - // Handle email-style format: username@domain.com - if idx := strings.Index(bareUsername, "@"); idx != -1 { - bareUsername = bareUsername[:idx] - } - - return isWindowsPrivilegedUser(bareUsername) -} - -// isWindowsPrivilegedUser checks if a bare username (domain already stripped) represents a Windows privileged account -func isWindowsPrivilegedUser(bareUsername string) bool { - // common privileged usernames (case insensitive) - privilegedNames := []string{ - "administrator", - "admin", - "root", - "system", - "localsystem", - "networkservice", - "localservice", - } - - usernameLower := strings.ToLower(bareUsername) - for _, privilegedName := range privilegedNames { - if usernameLower == privilegedName { - return true - } - } - - // computer accounts (ending with $) are not privileged by themselves - // They only gain privileges through group membership or specific SIDs - - if targetUser, err := lookupUser(bareUsername); err == nil { - return isWindowsPrivilegedSID(targetUser.Uid) - } - - return false -} - -// isWindowsPrivilegedSID checks if a Windows SID represents a privileged account -func isWindowsPrivilegedSID(sid string) bool { - privilegedSIDs := []string{ - "S-1-5-18", // Local System (SYSTEM) - "S-1-5-19", // Local Service (NT AUTHORITY\LOCAL SERVICE) - "S-1-5-20", // Network Service (NT AUTHORITY\NETWORK SERVICE) - "S-1-5-32-544", // Administrators group (BUILTIN\Administrators) - "S-1-5-500", // Built-in Administrator account (local machine RID 500) - } - - for _, privilegedSID := range privilegedSIDs { - if sid == privilegedSID { - return true - } - } - - // Check for domain administrator accounts (RID 500 in any domain) - // Format: S-1-5-21-domain-domain-domain-500 - // This is reliable as RID 500 is reserved for the domain Administrator account - if strings.HasPrefix(sid, "S-1-5-21-") && strings.HasSuffix(sid, "-500") { - return true - } - - // Check for other well-known privileged RIDs in domain contexts - // RID 512 = Domain Admins group, RID 516 = Domain Controllers group - if strings.HasPrefix(sid, "S-1-5-21-") { - if strings.HasSuffix(sid, "-512") || // Domain Admins group - strings.HasSuffix(sid, "-516") || // Domain Controllers group - strings.HasSuffix(sid, "-519") { // Enterprise Admins group - return true - } - } - - return false + return getWindowsAccountPrivilegedOrUnknown(username) } // isCurrentProcessPrivileged checks if the current process is running with elevated privileges. // On Unix systems, this means running as root (UID 0). -// On Windows, this means running as Administrator or SYSTEM. +// On Windows, this means the process token is elevated (administrators, SYSTEM). func isCurrentProcessPrivileged() bool { if getCurrentOS() == "windows" { - return isWindowsElevated() + return getProcessElevated() } return getEuid() == 0 } - -// isWindowsElevated checks if the current process is running with elevated privileges on Windows -func isWindowsElevated() bool { - currentUser, err := getCurrentUser() - if err != nil { - log.Errorf("failed to get current user for privilege check, assuming non-privileged: %v", err) - return false - } - - if isWindowsPrivilegedSID(currentUser.Uid) { - log.Debugf("Windows user switching supported: running as privileged SID %s", currentUser.Uid) - return true - } - - if isPrivilegedUsername(currentUser.Username) { - log.Debugf("Windows user switching supported: running as privileged username %s", currentUser.Username) - return true - } - - log.Debugf("Windows user switching not supported: not running as privileged user (current: %s)", currentUser.Uid) - return false -} diff --git a/client/ssh/server/user_utils_test.go b/client/ssh/server/user_utils_test.go index 637dc10d0..2fa9b68ee 100644 --- a/client/ssh/server/user_utils_test.go +++ b/client/ssh/server/user_utils_test.go @@ -4,6 +4,7 @@ import ( "errors" "os/user" "runtime" + "strings" "testing" "github.com/stretchr/testify/assert" @@ -27,8 +28,8 @@ func setupTestDependencies(currentUser *user.User, currentUserErr error, os stri originalLookupUser := lookupUser originalGetCurrentOS := getCurrentOS originalGetEuid := getEuid - - // Reset caches to ensure clean test state + originalGetProcessElevated := getProcessElevated + originalGetWindowsAccountPrivilegedOrUnknown := getWindowsAccountPrivilegedOrUnknown // Set test values - inject platform dependencies getCurrentUser = func() (*user.User, error) { @@ -53,16 +54,31 @@ func setupTestDependencies(currentUser *user.User, currentUserErr error, os stri return euid } - // Mock privilege detection based on the test user - getIsProcessPrivileged = func() bool { + // Simulate the Windows token elevation check based on the fixture user: + // the built-in Administrator (RID 500) and SYSTEM run elevated. + getProcessElevated = func() bool { if currentUser == nil { return false } - // Check both username and SID for Windows systems - if os == "windows" && isWindowsPrivilegedSID(currentUser.Uid) { + return currentUser.Uid == "S-1-5-18" || strings.HasSuffix(currentUser.Uid, "-500") + } + + // Simulate the Windows account classifier for the fixture accounts. + // "root" does not exist on Windows; the real classifier fails closed on + // unresolvable accounts, so it counts as privileged here too. + getWindowsAccountPrivilegedOrUnknown = func(username string) bool { + bare := username + if idx := strings.LastIndex(bare, `\`); idx != -1 { + bare = bare[idx+1:] + } + if idx := strings.Index(bare, "@"); idx != -1 { + bare = bare[:idx] + } + switch strings.ToLower(bare) { + case "administrator", "system", "root": return true } - return isPrivilegedUsername(currentUser.Username) + return false } // Return cleanup function @@ -71,10 +87,8 @@ func setupTestDependencies(currentUser *user.User, currentUserErr error, os stri lookupUser = originalLookupUser getCurrentOS = originalGetCurrentOS getEuid = originalGetEuid - - getIsProcessPrivileged = isCurrentProcessPrivileged - - // Reset caches after test + getProcessElevated = originalGetProcessElevated + getWindowsAccountPrivilegedOrUnknown = originalGetWindowsAccountPrivilegedOrUnknown } } @@ -421,6 +435,9 @@ func TestUsedFallback_MeansNoPrivilegeDropping(t *testing.T) { } func TestPrivilegedUsernameDetection(t *testing.T) { + // Windows classification is syscall-backed (SID resolution, group + // membership) and is covered by privileges_windows_test.go; here only the + // Unix logic and the platform dispatch are exercised. tests := []struct { name string username string @@ -432,25 +449,9 @@ func TestPrivilegedUsernameDetection(t *testing.T) { {"unix_regular_user", "alice", "linux", false}, {"unix_root_capital", "Root", "linux", false}, // Case-sensitive - // Windows tests + // Windows dispatch to the (mocked) account classifier {"windows_administrator", "Administrator", "windows", true}, - {"windows_system", "SYSTEM", "windows", true}, - {"windows_admin", "admin", "windows", true}, - {"windows_admin_lowercase", "administrator", "windows", true}, // Case-insensitive - {"windows_domain_admin", "DOMAIN\\Administrator", "windows", true}, - {"windows_email_admin", "admin@domain.com", "windows", true}, {"windows_regular_user", "alice", "windows", false}, - {"windows_domain_user", "DOMAIN\\alice", "windows", false}, - {"windows_localsystem", "localsystem", "windows", true}, - {"windows_networkservice", "networkservice", "windows", true}, - {"windows_localservice", "localservice", "windows", true}, - - // Computer accounts (these depend on current user context in real implementation) - {"windows_computer_account", "WIN2K19-C2$", "windows", false}, // Computer account by itself not privileged - {"windows_domain_computer", "DOMAIN\\COMPUTER$", "windows", false}, // Domain computer account - - // Cross-platform - {"root_on_windows", "root", "windows", true}, // Root should be privileged everywhere } for _, tt := range tests { @@ -459,50 +460,8 @@ func TestPrivilegedUsernameDetection(t *testing.T) { cleanup := setupTestDependencies(nil, nil, tt.platform, 1000, nil, nil) defer cleanup() - result := isPrivilegedUsername(tt.username) - assert.Equal(t, tt.privileged, result) - }) - } -} - -func TestWindowsPrivilegedSIDDetection(t *testing.T) { - tests := []struct { - name string - sid string - privileged bool - description string - }{ - // Well-known system accounts - {"system_account", "S-1-5-18", true, "Local System (SYSTEM)"}, - {"local_service", "S-1-5-19", true, "Local Service"}, - {"network_service", "S-1-5-20", true, "Network Service"}, - {"administrators_group", "S-1-5-32-544", true, "Administrators group"}, - {"builtin_administrator", "S-1-5-500", true, "Built-in Administrator"}, - - // Domain accounts - {"domain_administrator", "S-1-5-21-1234567890-1234567890-1234567890-500", true, "Domain Administrator (RID 500)"}, - {"domain_admins_group", "S-1-5-21-1234567890-1234567890-1234567890-512", true, "Domain Admins group"}, - {"domain_controllers_group", "S-1-5-21-1234567890-1234567890-1234567890-516", true, "Domain Controllers group"}, - {"enterprise_admins_group", "S-1-5-21-1234567890-1234567890-1234567890-519", true, "Enterprise Admins group"}, - - // Regular users - {"regular_user", "S-1-5-21-1234567890-1234567890-1234567890-1001", false, "Regular domain user"}, - {"another_regular_user", "S-1-5-21-1234567890-1234567890-1234567890-1234", false, "Another regular user"}, - {"local_user", "S-1-5-21-1234567890-1234567890-1234567890-1000", false, "Local regular user"}, - - // Groups that are not privileged - {"domain_users", "S-1-5-21-1234567890-1234567890-1234567890-513", false, "Domain Users group"}, - {"power_users", "S-1-5-32-547", false, "Power Users group"}, - - // Invalid SIDs - {"malformed_sid", "S-1-5-invalid", false, "Malformed SID"}, - {"empty_sid", "", false, "Empty SID"}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - result := isWindowsPrivilegedSID(tt.sid) - assert.Equal(t, tt.privileged, result, "Failed for %s: %s", tt.description, tt.sid) + result := isPrivilegedOrUnknown(tt.username) + assert.Equal(t, tt.privileged, result, "privilege classification for %s on %s", tt.username, tt.platform) }) } } diff --git a/client/ssh/server/userswitching_windows.go b/client/ssh/server/userswitching_windows.go index 260e1301e..9e8cd5b30 100644 --- a/client/ssh/server/userswitching_windows.go +++ b/client/ssh/server/userswitching_windows.go @@ -91,7 +91,7 @@ func validateUsernameFormat(username string) error { func (s *Server) createExecutorCommand(logger *log.Entry, session ssh.Session, localUser *user.User, hasPty bool) (*exec.Cmd, func(), error) { logger.Debugf("creating Windows executor command for user %s (Pty: %v)", localUser.Username, hasPty) - username, _ := s.parseUsername(localUser.Username) + username, _ := parseUsername(localUser.Username) if err := validateUsername(username); err != nil { return nil, nil, fmt.Errorf("invalid username %q: %w", username, err) } @@ -102,7 +102,7 @@ func (s *Server) createExecutorCommand(logger *log.Entry, session ssh.Session, l // createUserSwitchCommand creates a command with Windows user switching. // Returns the command and a cleanup function that must be called after starting the process. func (s *Server) createUserSwitchCommand(logger *log.Entry, session ssh.Session, localUser *user.User) (*exec.Cmd, func(), error) { - username, domain := s.parseUsername(localUser.Username) + username, domain := parseUsername(localUser.Username) shell := getUserShell(localUser.Uid) @@ -138,7 +138,7 @@ func (s *Server) createUserSwitchCommand(logger *log.Entry, session ssh.Session, } // parseUsername extracts username and domain from a Windows username -func (s *Server) parseUsername(fullUsername string) (username, domain string) { +func parseUsername(fullUsername string) (username, domain string) { // Handle DOMAIN\username format if idx := strings.LastIndex(fullUsername, `\`); idx != -1 { domain = fullUsername[:idx] diff --git a/client/ssh/session.go b/client/ssh/session.go new file mode 100644 index 000000000..999b6f251 --- /dev/null +++ b/client/ssh/session.go @@ -0,0 +1,84 @@ +package ssh + +import ( + "fmt" + "io" + + log "github.com/sirupsen/logrus" + "golang.org/x/crypto/ssh" +) + +// DefaultTerminalModes are the PTY modes used by the interactive terminal clients. +var DefaultTerminalModes = ssh.TerminalModes{ + ssh.ECHO: 1, + ssh.TTY_OP_ISPEED: 14400, + ssh.TTY_OP_OSPEED: 14400, + ssh.VINTR: 3, // Ctrl+C + ssh.VQUIT: 28, // Ctrl+\ + ssh.VERASE: 127, // Backspace + ssh.VKILL: 21, // Ctrl+U + ssh.VEOF: 4, // Ctrl+D + ssh.VEOL: 0, + ssh.VEOL2: 0, + ssh.VSTART: 17, // Ctrl+Q + ssh.VSTOP: 19, // Ctrl+S + ssh.VSUSP: 26, // Ctrl+Z + ssh.VDISCARD: 15, // Ctrl+O + ssh.VREPRINT: 18, // Ctrl+R + ssh.VWERASE: 23, // Ctrl+W + ssh.VLNEXT: 22, // Ctrl+V +} + +// PTYSession is an interactive shell session with a PTY and its I/O pipes. +type PTYSession struct { + Session *ssh.Session + Stdin io.WriteCloser + Stdout io.Reader + Stderr io.Reader +} + +// StartPTYSession opens a session on the client, requests an xterm-256color PTY +// with the default terminal modes, wires up the I/O pipes and starts a shell. +// The session is closed on any error. +func StartPTYSession(client *ssh.Client, cols, rows int) (*PTYSession, error) { + session, err := client.NewSession() + if err != nil { + return nil, fmt.Errorf("new session: %w", err) + } + + pty, err := setupPTYSession(session, cols, rows) + if err != nil { + if closeErr := session.Close(); closeErr != nil { + log.Debugf("ssh: session close after setup error: %v", closeErr) + } + return nil, err + } + return pty, nil +} + +// setupPTYSession requests the PTY, opens the pipes and starts the shell on an +// already created session. +func setupPTYSession(session *ssh.Session, cols, rows int) (*PTYSession, error) { + if err := session.RequestPty("xterm-256color", rows, cols, DefaultTerminalModes); err != nil { + return nil, fmt.Errorf("request pty: %w", err) + } + + stdin, err := session.StdinPipe() + if err != nil { + return nil, fmt.Errorf("stdin pipe: %w", err) + } + stdout, err := session.StdoutPipe() + if err != nil { + return nil, fmt.Errorf("stdout pipe: %w", err) + } + stderr, err := session.StderrPipe() + if err != nil { + return nil, fmt.Errorf("stderr pipe: %w", err) + } + + if err := session.Shell(); err != nil { + return nil, fmt.Errorf("start shell: %w", err) + } + + return &PTYSession{Session: session, Stdin: stdin, Stdout: stdout, Stderr: stderr}, nil +} diff --git a/client/status/status.go b/client/status/status.go index e8276d0fa..1c204cdb1 100644 --- a/client/status/status.go +++ b/client/status/status.go @@ -46,7 +46,10 @@ func ParseDaemonStatus(s string) DaemonStatus { // ConvertOptions holds parameters for ConvertToStatusOutputOverview. type ConvertOptions struct { - Anonymize bool + Anonymize bool + // AnonymizeLevel selects how much the anonymizer redacts. Only + // meaningful when Anonymize is set. + AnonymizeLevel anonymize.Level DaemonVersion string DaemonStatus DaemonStatus StatusFilter string @@ -217,6 +220,7 @@ func ConvertToStatusOutputOverview(pbFullStatus *proto.FullStatus, opts ConvertO if opts.Anonymize { anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses()) + anonymizer.SetLevel(opts.AnonymizeLevel) anonymizeOverview(anonymizer, &overview) } @@ -976,6 +980,7 @@ func timeAgo(t time.Time) string { func anonymizePeerDetail(a *anonymize.Anonymizer, peer *PeerStateDetailOutput) { peer.FQDN = a.AnonymizeDomain(peer.FQDN) + peer.PubKey = a.AnonymizeWGKey(peer.PubKey) if localIP, port, err := net.SplitHostPort(peer.IceCandidateEndpoint.Local); err == nil { peer.IceCandidateEndpoint.Local = fmt.Sprintf("%s:%s", a.AnonymizeIPString(localIP), port) } @@ -1007,6 +1012,7 @@ func anonymizeOverview(a *anonymize.Anonymizer, overview *OutputOverview) { overview.SignalState.URL = a.AnonymizeURI(overview.SignalState.URL) overview.SignalState.Error = a.AnonymizeString(overview.SignalState.Error) + overview.PubKey = a.AnonymizeWGKey(overview.PubKey) overview.IP = a.AnonymizeIPString(overview.IP) overview.IPv6 = a.AnonymizeIPString(overview.IPv6) for i, detail := range overview.Relays.Details { diff --git a/client/system/info.go b/client/system/info.go index 1838204b8..daeabca13 100644 --- a/client/system/info.go +++ b/client/system/info.go @@ -79,13 +79,15 @@ type Info struct { EnableSSHLocalPortForwarding bool EnableSSHRemotePortForwarding bool DisableSSHAuth bool + + SyncMessageVersion *int } func (i *Info) SetFlags( rosenpassEnabled, rosenpassPermissive bool, serverSSHAllowed *bool, disableClientRoutes, disableServerRoutes, - disableDNS, disableFirewall, blockLANAccess, blockInbound, disableIPv6 bool, + disableDNS, disableFirewall, blockLANAccess, blockInbound, disableIPv6 bool, syncMessageVersion *int, enableSSHRoot, enableSSHSFTP, enableSSHLocalPortForwarding, enableSSHRemotePortForwarding *bool, disableSSHAuth *bool, ) { @@ -103,6 +105,8 @@ func (i *Info) SetFlags( i.BlockInbound = blockInbound i.DisableIPv6 = disableIPv6 + i.SyncMessageVersion = syncMessageVersion + if enableSSHRoot != nil { i.EnableSSHRoot = *enableSSHRoot } diff --git a/client/system/info_android.go b/client/system/info_android.go index 3c71573bb..d4f479386 100644 --- a/client/system/info_android.go +++ b/client/system/info_android.go @@ -30,6 +30,11 @@ func GetInfo(ctx context.Context) *Info { kernelVersion = osInfo[2] } + addrs, err := networkAddresses() + if err != nil { + log.Warnf("discover network addresses: %s", err) + } + gio := &Info{ GoOS: runtime.GOOS, Kernel: kernel, @@ -41,6 +46,7 @@ func GetInfo(ctx context.Context) *Info { NetbirdVersion: version.NetbirdVersion(), UIVersion: extractUIVersion(ctx), KernelVersion: kernelVersion, + NetworkAddresses: addrs, SystemSerialNumber: serial(), SystemProductName: productModel(), SystemManufacturer: productManufacturer(), diff --git a/client/system/info_js.go b/client/system/info_js.go index f32532881..3323fb542 100644 --- a/client/system/info_js.go +++ b/client/system/info_js.go @@ -15,7 +15,7 @@ func UpdateStaticInfoAsync() { } // GetInfo retrieves system information for WASM environment -func GetInfo(_ context.Context) *Info { +func GetInfo(ctx context.Context) *Info { info := &Info{ GoOS: runtime.GOOS, Kernel: runtime.GOARCH, @@ -30,6 +30,13 @@ func GetInfo(_ context.Context) *Info { collectBrowserInfo(info) collectLocationInfo(info) collectSystemInfo(info) + + // A caller-provided device name wins, as on the other platforms. A peer + // registered over an API keeps reporting the name it was registered with, + // so its meta does not change on the first sync. + if name := extractDeviceName(ctx, info.Hostname); name != "" { + info.Hostname = name + } return info } diff --git a/client/system/info_js_test.go b/client/system/info_js_test.go new file mode 100644 index 000000000..e2a33ada0 --- /dev/null +++ b/client/system/info_js_test.go @@ -0,0 +1,27 @@ +//go:build js + +package system + +import ( + "context" + "testing" +) + +// TestGetInfoHonorsDeviceName covers a caller-provided device name reaching the +// reported hostname, so a peer registered over an API keeps reporting the name +// it was registered with instead of renaming itself on its first sync. +func TestGetInfoHonorsDeviceName(t *testing.T) { + ctx := context.WithValue(context.Background(), DeviceNameCtxKey, "session-name") + if got := GetInfo(ctx).Hostname; got != "session-name" { + t.Errorf("hostname should carry the caller's device name, got %q", got) + } +} + +// TestGetInfoWithoutDeviceNameKeepsFallback covers the embed layer's habit of +// always setting the context value: an empty name must not blank the hostname. +func TestGetInfoWithoutDeviceNameKeepsFallback(t *testing.T) { + ctx := context.WithValue(context.Background(), DeviceNameCtxKey, "") + if got := GetInfo(ctx).Hostname; got == "" { + t.Error("an empty device name must not blank the hostname") + } +} diff --git a/client/system/network_addr.go b/client/system/network_addr.go index 44260a938..505a6f0ea 100644 --- a/client/system/network_addr.go +++ b/client/system/network_addr.go @@ -1,4 +1,4 @@ -//go:build !ios +//go:build !ios && !android package system diff --git a/client/system/network_addr_android.go b/client/system/network_addr_android.go new file mode 100644 index 000000000..99a71e105 --- /dev/null +++ b/client/system/network_addr_android.go @@ -0,0 +1,89 @@ +package system + +import ( + "net/netip" + "strings" +) + +var iFaceDiscover IFaceDiscover + +type IFaceDiscover interface { + IFaces() (string, error) +} + +// SetIFaceDiscover configures the Android interface discovery provider. +func SetIFaceDiscover(discover IFaceDiscover) { + iFaceDiscover = discover +} + +func networkAddresses() ([]NetworkAddress, error) { + if iFaceDiscover == nil { + return nil, nil + } + ifaces, err := iFaceDiscover.IFaces() + if err != nil { + return nil, err + } + + var netAddresses []NetworkAddress + for _, line := range strings.Split(ifaces, "\n") { + addresses, ok := interfaceAddresses(line) + if !ok { + continue + } + for _, address := range addresses { + netAddr, ok := toNetworkAddress(address) + if !ok { + continue + } + if isDuplicated(netAddresses, netAddr) { + continue + } + netAddresses = append(netAddresses, netAddr) + } + } + return netAddresses, nil +} + +func interfaceAddresses(line string) ([]string, bool) { + parts := strings.Split(line, "|") + if len(parts) != 2 { + return nil, false + } + flags := strings.Fields(parts[0]) + if len(flags) != 8 { + return nil, false + } + up, loopback := flags[3], flags[5] + if up != "true" || loopback == "true" { + return nil, false + } + return strings.Fields(parts[1]), true +} + +func toNetworkAddress(address string) (NetworkAddress, bool) { + prefix, err := netip.ParsePrefix(address) + if err != nil { + return NetworkAddress{}, false + } + if prefix.Addr().Is4In6() { + if prefix.Bits() < 96 { + return NetworkAddress{}, false + } + prefix = netip.PrefixFrom(prefix.Addr().Unmap(), prefix.Bits()-96) + } + ip := prefix.Addr() + if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsMulticast() { + return NetworkAddress{}, false + } + return NetworkAddress{NetIP: prefix}, true +} + +func isDuplicated(addresses []NetworkAddress, addr NetworkAddress) bool { + for _, duplicated := range addresses { + if duplicated.NetIP == addr.NetIP { + return true + } + } + return false +} diff --git a/client/system/network_addr_test.go b/client/system/network_addr_test.go index a5f9c4279..b0be40f0a 100644 --- a/client/system/network_addr_test.go +++ b/client/system/network_addr_test.go @@ -1,4 +1,4 @@ -//go:build !ios +//go:build !ios && !android package system diff --git a/client/system/process_test.go b/client/system/process_test.go index 9d0a6b935..de1cfc1db 100644 --- a/client/system/process_test.go +++ b/client/system/process_test.go @@ -1,3 +1,5 @@ +//go:build windows || (linux && !android) || (darwin && !ios) || freebsd + package system import ( diff --git a/client/ui/authsession/service.go b/client/ui/authsession/service.go index 28efe7cfd..9c094c2a1 100644 --- a/client/ui/authsession/service.go +++ b/client/ui/authsession/service.go @@ -6,9 +6,11 @@ import ( "context" "time" + log "github.com/sirupsen/logrus" "google.golang.org/grpc/codes" gstatus "google.golang.org/grpc/status" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/proto" ) @@ -58,10 +60,21 @@ func (s *Session) RequestExtend(ctx context.Context, p ExtendStartParams) (Exten return ExtendStartResult{}, err } - req := &proto.RequestExtendAuthSessionRequest{} - if p.Hint != "" { - h := p.Hint - req.Hint = &h + // a request from the UI implies a graphical session, which the daemon cannot detect itself + req := &proto.RequestExtendAuthSessionRequest{HasGraphicalSession: true} + hint := p.Hint + if hint == "" { + pm := profilemanager.NewProfileManager() + if active, perr := pm.GetActiveProfile(); perr != nil { + log.Debugf("failed to get active profile for login hint: %v", perr) + } else if state, serr := pm.GetProfileState(active.ID); serr != nil { + log.Debugf("failed to get profile state for login hint: %v", serr) + } else { + hint = state.Email + } + } + if hint != "" { + req.Hint = &hint } resp, err := cli.RequestExtendAuthSession(ctx, req) diff --git a/client/ui/autostart_default.go b/client/ui/autostart_default.go index bf1b16a97..162922579 100644 --- a/client/ui/autostart_default.go +++ b/client/ui/autostart_default.go @@ -51,7 +51,7 @@ func autostartDisabledByMDM(policy *mdm.Policy) bool { // netbirdFootprintExists reports whether the machine already carries NetBird // daemon config or state, meaning this is not a genuinely fresh install. It is // the update-safety gate for the autostart default: upgrading users always -// have a footprint, so an update can never trigger a login-item write. +// have a footprint, so an update can never trigger a autostart entry write. func netbirdFootprintExists() bool { candidates := []string{ profilemanager.DefaultConfigPath, @@ -69,9 +69,23 @@ func netbirdFootprintExists() bool { // applyAutostartDefault runs the one-time launch-on-login default for genuinely // fresh installs. The autostartInitialized marker is persisted before any // enable attempt so a crash mid-flow degrades to "never enabled" instead of -// retrying login-item writes on every launch. A user's later disable in +// retrying autostart entry writes on every launch. A user's later disable in // Settings is never overridden: the marker guarantees at-most-once, ever. func applyAutostartDefault(ctx context.Context, autostart *services.Autostart, prefs *preferences.Store, prefsFileExisted bool) { + mdmDisabled := autostartDisabledByMDM(mdm.LoadPolicy()) + + if mdmDisabled { + if enabled, err := autostart.IsEnabled(ctx); err != nil { + log.Warnf("MDM disableAutostart: read autostart state: %v", err) + } else if enabled { + if err := autostart.SetEnabled(ctx, false); err != nil { + log.Warnf("MDM disableAutostart: force off failed: %v", err) + } else { + log.Info("MDM disableAutostart enforced: autostart turned off") + } + } + } + priorFootprint := netbirdFootprintExists() || prefsFileExisted if prefs.Get().AutostartInitialized { @@ -84,7 +98,7 @@ func applyAutostartDefault(ctx context.Context, autostart *services.Autostart, p state := autostartDefaultState{ supported: autostart.Supported(ctx), - mdmDisabled: autostartDisabledByMDM(mdm.LoadPolicy()), + mdmDisabled: mdmDisabled, priorInstall: priorFootprint, } enable, reason := shouldEnableAutostartDefault(state) diff --git a/client/ui/build/linux/nfpm/nfpm.yaml b/client/ui/build/linux/nfpm/nfpm.yaml index a05daef62..764855a63 100644 --- a/client/ui/build/linux/nfpm/nfpm.yaml +++ b/client/ui/build/linux/nfpm/nfpm.yaml @@ -26,17 +26,17 @@ contents: # Default dependencies for the GTK4 + WebKitGTK 6.0 stack (Ubuntu 24.04+ / Debian 13+) depends: - - libgtk-4-1 + - libgtk-4-1 (>= 4.14) - libwebkitgtk-6.0-4 - xdg-utils # Distribution-specific overrides for different package formats overrides: - # RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux + # RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux / openSUSE rpm: depends: - - gtk4 - - webkitgtk6.0 + - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) + - (webkitgtk6.0 or libwebkitgtk-6_0-4) - xdg-utils # Arch Linux packages diff --git a/client/ui/frontend/package.json b/client/ui/frontend/package.json index 3131b36cd..dcef99ad3 100644 --- a/client/ui/frontend/package.json +++ b/client/ui/frontend/package.json @@ -15,7 +15,8 @@ "lint": "eslint \"src/**/*.{ts,tsx}\"", "lint:fix": "eslint \"src/**/*.{ts,tsx}\" --fix", "check": "pnpm lint && pnpm typecheck && pnpm format:check", - "check:fix": "pnpm lint:fix && pnpm format && pnpm typecheck" + "check:fix": "pnpm lint:fix && pnpm format && pnpm typecheck", + "i18n:check": "node ../i18n/check-translations.mjs" }, "dependencies": { "@radix-ui/react-dialog": "^1.1.15", diff --git a/client/ui/frontend/src/components/CopyToClipboard.tsx b/client/ui/frontend/src/components/CopyToClipboard.tsx index 1b2a87da4..3cf681a1c 100644 --- a/client/ui/frontend/src/components/CopyToClipboard.tsx +++ b/client/ui/frontend/src/components/CopyToClipboard.tsx @@ -18,6 +18,9 @@ type CopyToClipboardProps = { className?: string; iconClassName?: string; alwaysShowIcon?: boolean; + // wrap lets long content (a shell command, a path) break across lines + // instead of being truncated to one line. + wrap?: boolean; variant?: CopyToClipboardVariant; "aria-label"?: string; tabIndex?: number; @@ -32,6 +35,7 @@ export const CopyToClipboard = ({ className, iconClassName, alwaysShowIcon = false, + wrap = false, variant = "default", "aria-label": ariaLabel, tabIndex = 0, @@ -83,7 +87,8 @@ export const CopyToClipboard = ({ > { + const { isReady } = useStatus(); + const sent = useRef(false); + + useEffect(() => { + if (!isReady || sent.current) return; + sent.current = true; + void Events.Emit(EVENT_WINDOW_PAINTED); + }, [isReady]); + + return null; +}; diff --git a/client/ui/frontend/src/components/empty-state/DaemonOutdatedOverlay.tsx b/client/ui/frontend/src/components/empty-state/DaemonOutdatedOverlay.tsx index 4ee8c2740..e8e7108eb 100644 --- a/client/ui/frontend/src/components/empty-state/DaemonOutdatedOverlay.tsx +++ b/client/ui/frontend/src/components/empty-state/DaemonOutdatedOverlay.tsx @@ -1,10 +1,13 @@ +import { useEffect, useState } from "react"; import { useTranslation } from "react-i18next"; import { AlertTriangleIcon, DownloadIcon } from "lucide-react"; import { Browser } from "@wailsio/runtime"; +import { Version } from "@bindings/services"; import { Button } from "@/components/buttons/Button"; import { useStatus } from "@/contexts/StatusContext.tsx"; const RELEASES_URL = "https://github.com/netbirdio/netbird/releases/latest"; +const RC_RELEASES_URL = "https://pkgs.netbird.io/releases/rc"; function openUrl(url: string) { Browser.OpenURL(url).catch(() => globalThis.open(url, "_blank")); @@ -12,7 +15,26 @@ function openUrl(url: string) { export const DaemonOutdatedOverlay = () => { const { t } = useTranslation(); - const { isDaemonOutdated } = useStatus(); + const { status, isDaemonOutdated } = useStatus(); + + const [guiVersion, setGuiVersion] = useState("-"); + const clientVersion = status?.daemonVersion ?? "—"; + + const isRc = /-rc/i.test(guiVersion) || /-rc/i.test(clientVersion); + const downloadUrl = isRc ? RC_RELEASES_URL : RELEASES_URL; + + useEffect(() => { + if (!isDaemonOutdated) return; + let cancelled = false; + Version.GUI() + .then((v) => { + if (!cancelled) setGuiVersion(v); + }) + .catch((err) => console.error("[DaemonOutdatedOverlay] GUI version error", err)); + return () => { + cancelled = true; + }; + }, [isDaemonOutdated]); if (!isDaemonOutdated) return null; @@ -38,10 +60,37 @@ export const DaemonOutdatedOverlay = () => {

{t("daemon.outdated.description")}

+
+

+ {clientVersion === "development" ? ( + + {t("settings.about.clientName")}{" "} + + {t("settings.about.development")} + + + ) : ( + t("settings.about.client", { version: clientVersion }) + )} +

+

+ {guiVersion === "development" ? ( + + {t("settings.about.guiName")}{" "} + + {t("settings.about.development")} + + + ) : ( + t("settings.about.gui", { version: guiVersion }) + )} +

+
+
-
diff --git a/client/ui/frontend/src/contexts/DebugBundleContext.tsx b/client/ui/frontend/src/contexts/DebugBundleContext.tsx index a0a131fbf..5f2ed9041 100644 --- a/client/ui/frontend/src/contexts/DebugBundleContext.tsx +++ b/client/ui/frontend/src/contexts/DebugBundleContext.tsx @@ -71,10 +71,12 @@ type BundleOptions = { hasWindow: boolean; totalSec: number; uploadUrl: string; - anonymize: boolean; + anonymizeLevel: AnonymizeLevel; systemInfo: boolean; }; +export type AnonymizeLevel = "none" | "default" | "strict"; + const startCaptureBestEffort = async (totalSec: number, pcap: CaptureState) => { try { // Mirror the CLI's safety margin: window + 30s, server caps at 10m. @@ -187,7 +189,10 @@ const runBundleFlow = async ( if (opts.uploadUrl) setStage({ kind: "uploading" }); const result = await DebugSvc.Bundle({ - anonymize: opts.anonymize, + anonymize: opts.anonymizeLevel !== "none", + // The daemon only knows "default" and "strict"; "none" is expressed + // through the anonymize flag being off. + anonymizeLevel: opts.anonymizeLevel === "strict" ? "strict" : "default", systemInfo: opts.systemInfo, uploadUrl: opts.uploadUrl, logFileCount, @@ -198,7 +203,7 @@ const runBundleFlow = async ( }; const useDebugBundle = () => { - const [anonymize, setAnonymize] = useState(false); + const [anonymizeLevel, setAnonymizeLevel] = useState("none"); const [systemInfo, setSystemInfo] = useState(true); const [upload, setUpload] = useState(true); const [trace, setTrace] = useState(true); @@ -240,7 +245,7 @@ const useDebugBundle = () => { hasWindow: capture && totalSec > 0, totalSec, uploadUrl: upload ? NETBIRD_UPLOAD_URL : "", - anonymize, + anonymizeLevel, systemInfo, }; @@ -272,8 +277,8 @@ const useDebugBundle = () => { }; return { - anonymize, - setAnonymize, + anonymizeLevel, + setAnonymizeLevel, systemInfo, setSystemInfo, upload, diff --git a/client/ui/frontend/src/contexts/ProfileContext.tsx b/client/ui/frontend/src/contexts/ProfileContext.tsx index 4dd3eaa7a..62377f1bc 100644 --- a/client/ui/frontend/src/contexts/ProfileContext.tsx +++ b/client/ui/frontend/src/contexts/ProfileContext.tsx @@ -28,6 +28,7 @@ type ProfileContextValue = { loaded: boolean; refresh: () => Promise; switchProfile: (id: string) => Promise; + switchProfileNoConnect: (id: string) => Promise; addProfile: (name: string) => Promise; removeProfile: (id: string) => Promise; renameProfile: (id: string, newName: string) => Promise; @@ -112,6 +113,16 @@ export const ProfileProvider = ({ children }: { children: ReactNode }) => { [username, refresh], ); + // Manage-profiles variant: switches without connecting, so the user can + // still adjust the management URL before bringing the connection up. + const switchProfileNoConnect = useCallback( + async (id: string) => { + await ProfileSwitcher.SwitchActiveNoConnect({ profileName: id, username }); + await refresh(); + }, + [username, refresh], + ); + // addProfile creates a profile by display name and returns the // daemon-generated ID, so the caller can immediately address it by ID. const addProfile = useCallback( @@ -158,6 +169,7 @@ export const ProfileProvider = ({ children }: { children: ReactNode }) => { loaded, refresh, switchProfile, + switchProfileNoConnect, addProfile, removeProfile, renameProfile, @@ -171,6 +183,7 @@ export const ProfileProvider = ({ children }: { children: ReactNode }) => { loaded, refresh, switchProfile, + switchProfileNoConnect, addProfile, removeProfile, renameProfile, diff --git a/client/ui/frontend/src/contexts/SettingsContext.tsx b/client/ui/frontend/src/contexts/SettingsContext.tsx index 37627bc76..3f4b2d0d2 100644 --- a/client/ui/frontend/src/contexts/SettingsContext.tsx +++ b/client/ui/frontend/src/contexts/SettingsContext.tsx @@ -14,7 +14,7 @@ import type { Config } from "@bindings/services/models.js"; import i18next from "@/lib/i18n"; import { useProfile } from "@/contexts/ProfileContext.tsx"; import { SettingsSkeleton } from "@/modules/settings/SettingsSkeleton.tsx"; -import { errorDialog, formatErrorMessage as errorMessage } from "@/lib/errors.ts"; +import { errorCommand, errorDialog, formatErrorMessage as errorMessage } from "@/lib/errors.ts"; const SAVE_DEBOUNCE_MS = 400; @@ -68,6 +68,21 @@ const useSettingsState = () => { loadedRef.current = loaded; }, [loaded]); + // reload re-reads the daemon's config, which is authoritative. Used on + // mount, on the daemon's config_changed event, and to undo an optimistic + // update the daemon then rejected. + const reload = useCallback( + async (profileName: string) => { + try { + const data = await SettingsSvc.GetConfig({ profileName, username }); + setLoaded({ profileName, data }); + } catch (e) { + console.warn("[SettingsContext] reload after rejected save failed", e); + } + }, + [username], + ); + useEffect(() => { if (!profileLoaded || !activeProfileId) return; let cancelled = false; @@ -133,13 +148,20 @@ const useSettingsState = () => { username, }); } catch (e) { + // The optimistic update is wrong now: the daemon refused it + // (a change that needs elevated privileges, an MDM-managed + // field, ...). Snap the controls back to what it actually + // holds before reporting, so the UI never shows a value the + // daemon does not have. + await reload(profileName); await errorDialog({ Title: i18next.t("settings.error.saveTitle"), Message: errorMessage(e), + Command: errorCommand(e), }); } }, - [username], + [username, reload], ); const setField = useCallback( diff --git a/client/ui/frontend/src/hooks/useKeepConnectedOnQuit.ts b/client/ui/frontend/src/hooks/useKeepConnectedOnQuit.ts new file mode 100644 index 000000000..eb68dd997 --- /dev/null +++ b/client/ui/frontend/src/hooks/useKeepConnectedOnQuit.ts @@ -0,0 +1,35 @@ +import { useCallback, useEffect, useState } from "react"; +import { Preferences } from "@bindings/services"; + +export const useKeepConnectedOnQuit = () => { + const [keepConnected, setKeepConnected] = useState(null); + + useEffect(() => { + let cancelled = false; + Preferences.Get() + .then((prefs) => { + if (cancelled) return; + setKeepConnected(prefs?.keepConnectedOnQuit ?? false); + }) + .catch((err: unknown) => { + if (cancelled) return; + console.warn("[useKeepConnectedOnQuit] load preferences failed", err); + setKeepConnected(false); + }); + return () => { + cancelled = true; + }; + }, []); + + const setKeepConnectedOnQuit = useCallback(async (keep: boolean) => { + setKeepConnected(keep); + try { + await Preferences.SetKeepConnectedOnQuit(keep); + } catch (err: unknown) { + setKeepConnected(!keep); + console.error("[useKeepConnectedOnQuit] SetKeepConnectedOnQuit failed", err); + } + }, []); + + return { keepConnected, setKeepConnectedOnQuit }; +}; diff --git a/client/ui/frontend/src/hooks/usePrivilege.ts b/client/ui/frontend/src/hooks/usePrivilege.ts new file mode 100644 index 000000000..05e9a7ce0 --- /dev/null +++ b/client/ui/frontend/src/hooks/usePrivilege.ts @@ -0,0 +1,32 @@ +import { useEffect, useState } from "react"; +import { Settings as SettingsSvc } from "@bindings/services"; +import { Privilege } from "@bindings/services/models.js"; + +// usePrivilege reports whether this UI process may perform the changes the daemon +// restricts to root/administrator. It is answered in-process from our own token +// with the daemon's own rule, so there is no round-trip and it works while the +// daemon is down. +// +// null means "not known yet", which includes the read having failed. Callers must +// treat that as "do not restrict": the daemon enforces this regardless, so the +// only thing a wrong guess here costs is a control that looks unavailable when it +// is not, or a save that fails with the daemon's own guidance. +export const usePrivilege = (): Privilege | null => { + const [privilege, setPrivilege] = useState(null); + + useEffect(() => { + let cancelled = false; + SettingsSvc.Privilege() + .then((p) => { + if (!cancelled) setPrivilege(p); + }) + .catch((e: unknown) => { + console.warn("[usePrivilege] read failed, not restricting controls", e); + }); + return () => { + cancelled = true; + }; + }, []); + + return privilege; +}; diff --git a/client/ui/frontend/src/layouts/AppLayout.tsx b/client/ui/frontend/src/layouts/AppLayout.tsx index 1588d9d08..0c2837b53 100644 --- a/client/ui/frontend/src/layouts/AppLayout.tsx +++ b/client/ui/frontend/src/layouts/AppLayout.tsx @@ -5,6 +5,7 @@ import { DebugBundleProvider } from "@/contexts/DebugBundleContext.tsx"; import { ProfileProvider } from "@/contexts/ProfileContext.tsx"; import { DialogProvider } from "@/contexts/DialogContext.tsx"; import { RestrictionsProvider } from "@/contexts/RestrictionsContext.tsx"; +import { ReadySignal } from "@/components/ReadySignal.tsx"; export const AppLayout = () => { return ( @@ -16,6 +17,7 @@ export const AppLayout = () => { + diff --git a/client/ui/frontend/src/lib/connection.ts b/client/ui/frontend/src/lib/connection.ts index fca03fc87..cc0e67cb3 100644 --- a/client/ui/frontend/src/lib/connection.ts +++ b/client/ui/frontend/src/lib/connection.ts @@ -43,7 +43,12 @@ function buildSsoCancelPromise(state: SsoState, signal?: AbortSignal): Promise { @@ -56,7 +61,7 @@ async function runSsoLogin( // suspended, so a frontend-driven Up (a promise continuation) would not // fire until the user woke the window (e.g. hovering the tray icon). const waitPromise = Connection.WaitSSOLoginAndUp( - { userCode: result.userCode, hostname: "" }, + { userCode: result.userCode, hostname: "", profileId: result.profileId }, { profileName: "", username: "" }, ); diff --git a/client/ui/frontend/src/lib/errors.ts b/client/ui/frontend/src/lib/errors.ts index 34a90dace..b4dee2717 100644 --- a/client/ui/frontend/src/lib/errors.ts +++ b/client/ui/frontend/src/lib/errors.ts @@ -1,6 +1,6 @@ import { WindowManager } from "@bindings/services"; -type ClassifiedError = { short: string; long: string }; +type ClassifiedError = { short: string; long: string; command: string }; const asObject = (v: unknown): Record | null => v && typeof v === "object" ? (v as Record) : null; @@ -22,20 +22,24 @@ const toWailsEnvelope = (e: unknown): Record | null => { return asObject(obj.cause) ?? parseJsonObject(obj.message); }; -// Read { short, long } from wherever the classified error sits in the envelope +// Read { short, long, command } from wherever the classified error sits in the envelope const toClassifiedError = (v: unknown): ClassifiedError | null => { const o = asObject(v); if (!o) return null; const short = typeof o.short === "string" ? o.short : ""; const long = typeof o.long === "string" ? o.long : ""; - return short || long ? { short, long } : null; + const command = typeof o.command === "string" ? o.command : ""; + return short || long ? { short, long, command } : null; +}; + +const classify = (e: unknown): ClassifiedError | null => { + const envelope = toWailsEnvelope(e); + return toClassifiedError(envelope?.cause) ?? toClassifiedError(envelope); }; export const formatErrorMessage = (e: unknown): string => { - const envelope = toWailsEnvelope(e); - // Prefer the structured { short, long } the daemon classifier produced. - const classified = toClassifiedError(envelope?.cause) ?? toClassifiedError(envelope); + const classified = classify(e); if (classified) { const { short, long } = classified; if (short && long && long !== short) return `${short} Details: ${long}`; @@ -44,17 +48,26 @@ export const formatErrorMessage = (e: unknown): string => { } // Unclassified (a service returned the raw daemon error) + const envelope = toWailsEnvelope(e); const message = envelope?.message; if (typeof message === "string" && message) return message; if (e instanceof Error) return e.message; return String(e); }; +// errorCommand returns a command the user can run to complete an operation the +// daemon refused, when the error carries one (a change that needs elevated +// privileges). Empty for every other error. +export const errorCommand = (e: unknown): string => classify(e)?.command ?? ""; + export type ErrorDialogOptions = { Title: string; Message: string; + // Command is shown for copying below the message. Defaults to the one the + // error carries, so callers only pass it to override. + Command?: string; }; export function errorDialog(options: ErrorDialogOptions): Promise { - return WindowManager.OpenError(options.Title, options.Message); + return WindowManager.OpenError(options.Title, options.Message, options.Command ?? ""); } diff --git a/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx b/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx index 4fb5e2586..4861c492a 100644 --- a/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx +++ b/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx @@ -2,6 +2,7 @@ import { type ReactNode } from "react"; import { useTranslation } from "react-i18next"; import { Browser } from "@wailsio/runtime"; import { DownloadIcon, NotepadText } from "lucide-react"; +import { Update as UpdateSvc } from "@bindings/services"; import { Button } from "@/components/buttons/Button"; import { useClientVersion } from "@/contexts/ClientVersionContext"; import { cn } from "@/lib/cn"; @@ -14,6 +15,12 @@ function openUrl(url: string) { }); } +function openInstallerDownload() { + UpdateSvc.DownloadURL() + .then(openUrl) + .catch(() => openUrl(GITHUB_RELEASES)); +} + export function UpdateVersionCard() { const { t } = useTranslation(); const { updateVersion, enforced, triggerUpdate } = useClientVersion(); @@ -37,11 +44,7 @@ export function UpdateVersionCard() { {t("update.card.installNow")} ) : ( - diff --git a/client/ui/frontend/src/modules/error/ErrorDialog.tsx b/client/ui/frontend/src/modules/error/ErrorDialog.tsx index f7929fad2..4fbb78052 100644 --- a/client/ui/frontend/src/modules/error/ErrorDialog.tsx +++ b/client/ui/frontend/src/modules/error/ErrorDialog.tsx @@ -3,6 +3,7 @@ import { useTranslation } from "react-i18next"; import { useSearchParams } from "react-router-dom"; import { AlertCircleIcon } from "lucide-react"; import { Button } from "@/components/buttons/Button"; +import { CopyToClipboard } from "@/components/CopyToClipboard"; import { ConfirmDialog } from "@/components/dialog/ConfirmDialog"; import { DialogActions } from "@/components/dialog/DialogActions"; import { DialogDescription } from "@/components/dialog/DialogDescription"; @@ -12,14 +13,22 @@ import { WindowManager } from "@bindings/services"; import { useAutoSizeWindow } from "@/hooks/useAutoSizeWindow"; const WINDOW_WIDTH = 380; +// A command needs the room to wrap at a sensible number of characters instead of +// breaking every few words. +const WINDOW_WIDTH_WITH_COMMAND = 460; export default function ErrorDialog() { const { t } = useTranslation(); - const contentRef = useAutoSizeWindow(WINDOW_WIDTH); const [params] = useSearchParams(); const title = params.get("title") || t("window.title.error"); const message = params.get("message") || ""; + // Set when the daemon refused an operation that needs elevated privileges: + // the command that performs it, offered for copying. + const command = params.get("command") || ""; + const contentRef = useAutoSizeWindow( + command ? WINDOW_WIDTH_WITH_COMMAND : WINDOW_WIDTH, + ); const close = useCallback(() => { WindowManager.CloseError().catch(console.error); @@ -37,15 +46,37 @@ export default function ErrorDialog() { -
+
{title} {message && ( - {message} + {/* select-text: the message often names a path, a flag or an + address the user needs to act on. */} + + {message} + )} + {command && ( + + + {command} + + + )}
diff --git a/client/ui/frontend/src/modules/profiles/ProfilesTab.tsx b/client/ui/frontend/src/modules/profiles/ProfilesTab.tsx index c1ce2e449..97261ccc9 100644 --- a/client/ui/frontend/src/modules/profiles/ProfilesTab.tsx +++ b/client/ui/frontend/src/modules/profiles/ProfilesTab.tsx @@ -45,7 +45,7 @@ export function ProfilesTab() { activeProfileId, loaded, username, - switchProfile, + switchProfileNoConnect, addProfile, removeProfile, renameProfile, @@ -100,7 +100,7 @@ export function ProfilesTab() { confirmLabel: t("profile.switch.confirm"), }); if (!ok) return; - await guarded(i18next.t("profile.error.switchTitle"), () => switchProfile(id)); + await guarded(i18next.t("profile.error.switchTitle"), () => switchProfileNoConnect(id)); }; const handleDeregister = async (id: string, name: string) => { @@ -129,14 +129,13 @@ export function ProfilesTab() { await guarded(i18next.t("profile.error.createTitle"), async () => { const id = await addProfile(name); // SetConfig is keyed by the new profile's ID, so it writes the - // not-yet-active profile. Write before switching so any reconnect - // targets the right deployment. + // not-yet-active profile before the switch makes it current. if (!isNetbirdCloud(managementUrl)) { await SettingsSvc.SetConfig( new SetConfigParams({ profileName: id, username, managementUrl }), ); } - await switchProfile(id); + await switchProfileNoConnect(id); }); }; diff --git a/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx b/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx index 2ceb958d4..ef8d6862f 100644 --- a/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx +++ b/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx @@ -11,7 +11,7 @@ import { DialogHeading } from "@/components/dialog/DialogHeading"; import { SquareIcon } from "@/components/SquareIcon"; import { Connection, Profiles as ProfilesSvc, Session, WindowManager } from "@bindings/services"; import { useAutoSizeWindow } from "@/hooks/useAutoSizeWindow"; -import { EVENT_BROWSER_LOGIN_CANCEL } from "@/lib/connection"; +import { EVENT_BROWSER_LOGIN_CANCEL, EVENT_TRIGGER_LOGIN } from "@/lib/connection"; import { errorDialog, formatErrorMessage } from "@/lib/errors.ts"; import { formatRemaining } from "@/lib/formatters"; @@ -73,6 +73,13 @@ export default function SessionExpirationDialog() { let offCancel: (() => void) | undefined; + // Return the dialog to its interactive state and dismiss the browser popup + const resetDialog = () => { + offCancel?.(); + WindowManager.CloseBrowserLogin().catch(console.error); + setBusy(false); + }; + try { const start = await Session.RequestExtend({ hint: "" }); const uri = start.verificationUriComplete || start.verificationUri; @@ -105,25 +112,37 @@ export default function SessionExpirationDialog() { if (outcome.kind === "cancel") { waitPromise.cancel?.(); waitPromise.catch(() => {}); + resetDialog(); return; } // Another surface owns this flow; keep the dialog open to retry. if (outcome.result.preempted) { + resetDialog(); return; } - - // Close before the popup so the restore can't flash this window back. - WindowManager.CloseSessionExpiration().catch(console.error); + WindowManager.CloseRenewFlow().catch(console.error); } catch (e) { + resetDialog(); await errorDialog({ Title: t("sessionExpiration.extendFailedTitle"), Message: formatErrorMessage(e), }); - } finally { - offCancel?.(); - WindowManager.CloseBrowserLogin().catch(console.error); + } + }, [busy, t]); + + const authenticate = useCallback(async () => { + if (busy) return; + setBusy(true); + try { + await Events.Emit(EVENT_TRIGGER_LOGIN); + await WindowManager.CloseSessionExpiration(); + } catch (e) { setBusy(false); + await errorDialog({ + Title: t("connect.error.loginTitle"), + Message: formatErrorMessage(e), + }); } }, [busy, t]); @@ -139,12 +158,11 @@ export default function SessionExpirationDialog() { }); WindowManager.CloseSessionExpiration().catch(console.error); } catch (e) { + setBusy(false); await errorDialog({ Title: t("sessionExpiration.logoutFailedTitle"), Message: formatErrorMessage(e), }); - } finally { - setBusy(false); } }, [busy, t]); @@ -182,7 +200,7 @@ export default function SessionExpirationDialog() { variant={"primary"} size={"md"} className={"w-full"} - onClick={stay} + onClick={expired ? authenticate : stay} disabled={busy} > {expired ? t("sessionExpiration.authenticate") : t("sessionExpiration.stay")} diff --git a/client/ui/frontend/src/modules/settings/SettingsGeneral.tsx b/client/ui/frontend/src/modules/settings/SettingsGeneral.tsx index 05d40e15c..71720aebe 100644 --- a/client/ui/frontend/src/modules/settings/SettingsGeneral.tsx +++ b/client/ui/frontend/src/modules/settings/SettingsGeneral.tsx @@ -11,6 +11,7 @@ import { ManagementServerSwitch } from "@/components/ManagementServerSwitch.tsx" import { ManagementMode, useManagementUrl } from "@/hooks/useManagementUrl.ts"; import { LanguagePicker } from "@/components/LanguagePicker.tsx"; import { useRestrictions } from "@/contexts/RestrictionsContext.tsx"; +import { useKeepConnectedOnQuit } from "@/hooks/useKeepConnectedOnQuit.ts"; export function SettingsGeneral() { const { t } = useTranslation(); @@ -19,6 +20,7 @@ export function SettingsGeneral() { const { mode, setMode, setUrl, displayUrl, showError, canSave, save, checking, unreachable } = useManagementUrl(); const { mdm, features } = useRestrictions(); + const { keepConnected, setKeepConnectedOnQuit } = useKeepConnectedOnQuit(); const inputRef = useRef(null); const managementUrlId = useId(); @@ -57,6 +59,15 @@ export function SettingsGeneral() { helpText={t("settings.general.autostart.help")} /> )} + { + void setKeepConnectedOnQuit(v); + }} + loading={keepConnected === null} + label={t("settings.general.keepConnectedOnQuit.label")} + helpText={t("settings.general.keepConnectedOnQuit.help")} + /> {!mdm.managementURL && !features.disableUpdateSettings && ( diff --git a/client/ui/frontend/src/modules/settings/SettingsSSH.tsx b/client/ui/frontend/src/modules/settings/SettingsSSH.tsx index f18fd9493..bd91e520c 100644 --- a/client/ui/frontend/src/modules/settings/SettingsSSH.tsx +++ b/client/ui/frontend/src/modules/settings/SettingsSSH.tsx @@ -1,4 +1,5 @@ import { useTranslation } from "react-i18next"; +import { CopyToClipboard } from "@/components/CopyToClipboard"; import FancyToggleSwitch from "@/components/switches/FancyToggleSwitch"; import { HelpText } from "@/components/typography/HelpText"; import { Input } from "@/components/inputs/Input"; @@ -6,12 +7,50 @@ import { Label } from "@/components/typography/Label"; import { cn } from "@/lib/cn"; import { SectionGroup } from "@/modules/settings/SettingsSection.tsx"; import { useSettings } from "@/contexts/SettingsContext.tsx"; -import { type ChangeEvent, useEffect, useId, useState } from "react"; +import { usePrivilege } from "@/hooks/usePrivilege.ts"; +import { Privilege } from "@bindings/services/models.js"; +import { type ChangeEvent, type ReactNode, useEffect, useId, useState } from "react"; export function SettingsSSH() { const { t } = useTranslation(); const { config, setField } = useSettings(); + const privilege = usePrivilege(); const isSSHServerEnabled = config.serverSshAllowed; + + // The daemon restricts only the direction that hands out shells from a process + // running as root. So for an unprivileged user a guarded control is either + // unavailable (it is off and only they could turn it on) or a one-way switch + // (it is on, they may turn it off, but not back on) — say which, either way. + // + // A null privilege means we could not determine it: leave the control alone + // rather than greying it out with nothing to explain why. The daemon enforces + // this regardless, and a rejected save reports its own guidance. + const guarded = ( + guardedDirectionActive: boolean, + command: (p: Privilege) => string, + // inverted marks a control whose guarded direction is switching it off, so + // the one-way warning has to read the other way round. + inverted = false, + ) => { + if (!privilege || privilege.privileged) { + return { disabled: false, hint: undefined }; + } + const hint = ( + + ); + return { disabled: !guardedDirectionActive, hint }; + }; + + const sshServer = guarded(config.serverSshAllowed, (p) => p.allowSshServer); + const sshRoot = guarded(config.enableSshRoot, (p) => p.enableSshRoot); + // Inverted control: the guarded direction is switching authentication off, so + // it is the already-disabled state that is the one-way one. + const sshAuth = guarded(config.disableSshAuth, (p) => p.disableSshAuth, true); const jwtTtlId = useId(); const [jwtTtlInput, setJwtTtlInput] = useState(String(config.sshJwtCacheTtl)); @@ -46,9 +85,11 @@ export function SettingsSSH() { setField("serverSshAllowed", v)} + disabled={sshServer.disabled} label={t("settings.ssh.server.label")} helpText={t("settings.ssh.server.help")} /> + {sshServer.hint} setField("enableSshRoot", v)} + disabled={sshRoot.disabled} label={t("settings.ssh.root.label")} helpText={t("settings.ssh.root.help")} /> + {sshRoot.hint} setField("enableSshSftp", v)} @@ -88,9 +131,11 @@ export function SettingsSSH() { setField("disableSshAuth", !v)} + disabled={sshAuth.disabled} label={t("settings.ssh.jwt.label")} helpText={t("settings.ssh.jwt.help")} /> + {sshAuth.hint}
); } + +// PrivilegeHint explains what an unprivileged user can and cannot do with a +// guarded control, and offers the command that does it with the privileges the +// daemon requires. oneWay covers the control being in the guarded state already: +// switching it back is the part that needs privileges. +function PrivilegeHint({ + actor, + command, + oneWay, + inverted, +}: { + actor: string; + command: string; + oneWay: boolean; + inverted: boolean; +}): ReactNode { + const { t } = useTranslation(); + if (!command) return null; + return ( +
+ + {!oneWay + ? t("settings.ssh.privilege.hint", { actor }) + : inverted + ? t("settings.ssh.privilege.oneWayInverted", { actor }) + : t("settings.ssh.privilege.oneWay", { actor })} + + + + {command} + + +
+ ); +} diff --git a/client/ui/frontend/src/modules/settings/SettingsTroubleshooting.tsx b/client/ui/frontend/src/modules/settings/SettingsTroubleshooting.tsx index 991937719..8b1774ed6 100644 --- a/client/ui/frontend/src/modules/settings/SettingsTroubleshooting.tsx +++ b/client/ui/frontend/src/modules/settings/SettingsTroubleshooting.tsx @@ -1,6 +1,6 @@ import { useId, type ReactNode } from "react"; import { Trans, useTranslation } from "react-i18next"; -import { CircleCheckBig, FolderOpen, Loader2 } from "lucide-react"; +import { ChevronDown, CircleCheckBig, FolderOpen, Info, Loader2 } from "lucide-react"; import { Browser } from "@wailsio/runtime"; import { Debug as DebugSvc } from "@bindings/services"; import type { DebugBundleResult } from "@bindings/services/models.js"; @@ -8,13 +8,22 @@ import { Button } from "@/components/buttons/Button"; import { DialogActions } from "@/components/dialog/DialogActions"; import { DialogDescription } from "@/components/dialog/DialogDescription"; import { DialogHeading } from "@/components/dialog/DialogHeading"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuRadioGroup, + DropdownMenuRadioItem, + DropdownMenuTrigger, +} from "@/components/DropdownMenu"; import FancyToggleSwitch from "@/components/switches/FancyToggleSwitch"; import HelpText from "@/components/typography/HelpText.tsx"; import { Input } from "@/components/inputs/Input"; import { Label } from "@/components/typography/Label"; import { SquareIcon } from "@/components/SquareIcon"; +import { Tooltip } from "@/components/Tooltip"; +import { cn } from "@/lib/cn"; import { formatRemaining } from "@/lib/formatters"; -import type { DebugStage } from "@/contexts/DebugBundleContext"; +import type { AnonymizeLevel, DebugStage } from "@/contexts/DebugBundleContext"; import { useDebugBundleContext } from "@/contexts/DebugBundleContext"; import { SectionGroup, SettingsBottomBar } from "@/modules/settings/SettingsSection.tsx"; @@ -24,8 +33,8 @@ export function SettingsTroubleshooting() { const { t } = useTranslation(); const durationId = useId(); const { - anonymize, - setAnonymize, + anonymizeLevel, + setAnonymizeLevel, systemInfo, setSystemInfo, upload, @@ -55,12 +64,71 @@ export function SettingsTroubleshooting() { return ( - +
+
+
+ } + > + + + + + + {t("settings.troubleshooting.anonymize.help")} + +
+
+ + + + + + setAnonymizeLevel(v as AnonymizeLevel)} + > + + {t("settings.troubleshooting.anonymize.none")} + + + {t("settings.troubleshooting.anonymize.default")} + + + {t("settings.troubleshooting.anonymize.strict")} + + + + +
+
) { const { t } = useTranslation(); const trayScreenshot = trayScreenshotForOS(); + // macOS has no tray — the icon sits in the menu bar, so the copy says so. + const titleKey = isMacOS() ? "welcome.titleMac" : "welcome.title"; + const descriptionKey = isMacOS() ? "welcome.descriptionMac" : "welcome.description"; return ( <> @@ -36,9 +39,9 @@ export function WelcomeStepTray({ onContinue }: Readonly)
- {t("welcome.title")} + {t(titleKey)} - {t("welcome.description")} + {t(descriptionKey)}
diff --git a/client/ui/grpc.go b/client/ui/grpc.go index c8e3aed76..5450e136d 100644 --- a/client/ui/grpc.go +++ b/client/ui/grpc.go @@ -5,14 +5,13 @@ package main import ( "fmt" "runtime" - "strings" "sync" "time" "google.golang.org/grpc" "google.golang.org/grpc/backoff" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/ui/desktop" ) @@ -36,9 +35,10 @@ func (c *Conn) Client() (proto.DaemonServiceClient, error) { return c.client, nil } - cc, err := grpc.NewClient( - strings.TrimPrefix(c.addr, "tcp://"), - grpc.WithTransportCredentials(insecure.NewCredentials()), + // Lazy on purpose: grpc.NewClient does not connect here, so a daemon that + // is down surfaces as a per-RPC Unavailable instead of blocking the UI. + target, opts := daemonaddr.DialTarget(daemonaddr.ResolveDaemonAddr(c.addr)) + opts = append(opts, grpc.WithUserAgent(desktop.GetUIUserAgent()), // Cap reconnect backoff at 5s; gRPC's default 120s MaxDelay would // leave the UI waiting 30-60s to notice a freshly-started daemon. @@ -51,6 +51,8 @@ func (c *Conn) Client() (proto.DaemonServiceClient, error) { }, }), ) + + cc, err := grpc.NewClient(target, opts...) if err != nil { return nil, fmt.Errorf("dial daemon: %w", err) } @@ -58,10 +60,12 @@ func (c *Conn) Client() (proto.DaemonServiceClient, error) { return c.client, nil } -// DaemonAddr returns the default daemon gRPC address: a Unix socket on Linux/macOS, TCP loopback on Windows. +// DaemonAddr returns the default daemon gRPC address: a Unix socket on +// Linux/macOS, a named pipe on Windows. The pipe carries the caller's token, +// which loopback TCP does not, so the daemon can tell who is calling. func DaemonAddr() string { if runtime.GOOS == "windows" { - return "tcp://127.0.0.1:41731" + return daemonaddr.WindowsPipeAddr } return "unix:///var/run/netbird.sock" } diff --git a/client/ui/i18n/TRANSLATING.md b/client/ui/i18n/TRANSLATING.md index 88cbb8b11..803e093c8 100644 --- a/client/ui/i18n/TRANSLATING.md +++ b/client/ui/i18n/TRANSLATING.md @@ -2,9 +2,24 @@ A short brief for translating the desktop UI — for any translator, human or AI agent (*"you"* = whoever's translating). -**Drive an agent with:** *"Read `i18n/TRANSLATING.md` and translate the UI to Russian"* — or *"…and review the existing German translation."* +**Translations are managed on Crowdin: .** Join the project, pick your language, and translate in the editor. Each string carries a context note (the `description` from the source file) telling you what it is and where it shows up, and the project's glossary, style guide, and QA checks mirror this document. -> 💡 **The one habit that matters most:** read each key's `description` before translating it. Labels are terse and ambiguous on their own; the `description` tells you what the string is, where it shows up, what to keep verbatim, and what it actually means. +> 💡 **The one habit that matters most:** read each string's context before translating it. Labels are terse and ambiguous on their own; the context tells you what the string is, where it shows up, what to keep verbatim, and what it actually means. + +--- + +## How contributions flow + +```text +i18n/locales/en/common.json ──sync──▶ Crowdin ──service PR──▶ i18n/locales//common.json +``` + +- `i18n/locales/en/common.json` is the source of truth. New and changed strings sync to Crowdin automatically (see `crowdin.yml` in the repository root). +- Crowdin opens and updates a service pull request with the translated bundles, keeping the source's file shape and key order. Keys nobody has translated yet are left out of the export; the app falls back to English for them at runtime. Maintainers review and merge that PR. +- Don't hand-edit `i18n/locales//common.json` in your own PRs: the next sync would conflict with or overwrite your changes. Translate on Crowdin instead. +- Missing your language? Request it on the Crowdin project page or in a [GitHub discussion](https://github.com/netbirdio/netbird/discussions). When a language first ships, a maintainer adds its row to `i18n/locales/_index.json` with `code`, `displayName` (the native name), and `englishName`, which puts it in the app's language picker. + +**Prefer translating with an AI agent?** That still works: drive it with *"Read `i18n/TRANSLATING.md` and translate the UI to Russian"* as before, but deliver the result to Crowdin instead of a pull request. Download your language's file from the Crowdin editor, let the agent translate it, and upload it back (the editor's offline translation flow). Crowdin runs its QA checks on upload, and the next service PR carries the strings into the repo. --- @@ -30,25 +45,6 @@ A **business zero-trust VPN** — an encrypted **overlay mesh** between a compan --- -## The files - -``` -i18n/locales/_index.json shipped-language list -i18n/locales/en/common.json source of truth — message + description -i18n/locales//common.json a target — message only -``` - -Chrome-extension JSON, each key → `{ "message", "description" }`. You translate the **`message`**. - -| ✅ Do | ❌ Don't | -|---|---| -| Keep **every key** from `en`, in the same order | Translate, rename, reorder, drop, or add keys (they're identifiers; the set grows over time) | -| Put **only `message`** in target bundles | Copy `description` into a target bundle | -| Give every key a non-empty `message` | Leave keys missing or empty | -| Save valid UTF-8 JSON, no BOM | Add trailing commas or break the JSON | - ---- - ## Hard rules — get these exactly right These are the usual ways a translation *breaks the app*, not just reads oddly. @@ -58,7 +54,7 @@ These are the usual ways a translation *breaks the app*, not just reads oddly. | Copy `{placeholders}` verbatim — `{version}`, `{count}`, `{name}`… | Translate the word inside the braces (`{verbleibend}` breaks it) | | Reposition a placeholder so the sentence flows | Drop or duplicate a placeholder | | Preserve every `\n`, leading/trailing space, and trailing `...` | Trim "invisible" spaces or the `...` (they're load-bearing) | -| Keep `®` in WireGuard® and quotes around `{name}` | Strip punctuation the description flags | +| Keep `®` in WireGuard® and quotes around `{name}` | Strip punctuation the context flags | **Plurals:** the app has only a *one / other* split — the singular key fires only when `count == 1`; the `{count}` key covers everything else (0, 2, 5, 100…). Languages with more than two forms (ru, pl, uk) can't be fully correct here — use the form that fits the widest range (Russian genitive plural: `минут` / `часов` / `дней`). Don't invent extra keys or cram multiple forms into one string. When no single form fits every value — a unit label after a number field, say — reach for a number-agnostic form (an abbreviation, or wording that reads the same for 1 and 100) instead of forcing a plural the *one / other* split can't supply. @@ -78,13 +74,15 @@ When a brand sits beside a common noun, keep its exact spelling but join them th > **Use the word that language's IT users actually say.** Translate when a natural, common term exists; keep the English term *only* when the literal translation would be awkward or no one in that field really uses it. -Apply each term **consistently** — same English term → same translation everywhere — and keep a term once you've settled it. Whether a term stays English or takes a native word is **language-dependent**: a technical loanword (e.g. *Daemon*, *Handshake*) often stays, an everyday word (e.g. *Latency*, *Public key*) usually localizes, and some (*Exit Node*, *Peer*) go either way depending on the language. Decide per term with the rule above — a foreign origin alone is no reason to keep English. **Your main reference is the existing bundles:** match how a term was already rendered for your language rather than re-deciding it. +Apply each term **consistently** — same English term → same translation everywhere — and keep a term once you've settled it. Whether a term stays English or takes a native word is **language-dependent**: a technical loanword (e.g. *Daemon*, *Handshake*) often stays, an everyday word (e.g. *Latency*, *Public key*) usually localizes, and some (*Exit Node*, *Peer*) go either way depending on the language. Decide per term with the rule above — a foreign origin alone is no reason to keep English. **Your main reference is the existing translation:** match how a term was already rendered for your language rather than re-deciding it. Two checks before you commit a term: - **Prefer established localized wording.** If a widely used tool in this space (for example WireGuard) ships your language, its wording for a shared term such as *handshake* is what users already expect — look at the translated app, not just English docs. For generic UI verbs and formal address, follow your OS vendor's style guide (Microsoft / Apple / Google). - **Watch for false friends.** A literal translation can collide with a *different* established term in your field — confirm your word doesn't already mean something else in this domain before using it. +These tiers are mirrored in the Crowdin project glossary, so the editor highlights them inline. When you settle a new Tier C term for your language, add its translation to the glossary entry so it sticks for everyone who comes after you. + --- ## Style @@ -98,7 +96,7 @@ Two checks before you commit a term: Where it reads naturally, aim to keep each string **roughly the same length** as the English — the UI is tight and over-long strings can wrap or truncate. It's a soft preference, not a rule: if your language simply needs more words, use them. -A few habits that keep a bundle reading like one product rather than a word-for-word port: +A few habits that keep a translation reading like one product rather than a word-for-word port: - **Translate meaning, not words.** Render what a string *does*. An idiom or an awkward source phrase should become natural in your language, not a literal calque. - **Keep one voice within a family.** Sibling strings — the connection states, every settings *help* caption, every "… Failed" title — should share a grammatical form. If one member sounds wrong in that form, re-voice the whole family rather than leave one odd sibling. @@ -107,27 +105,26 @@ A few habits that keep a bundle reading like one product rather than a word-for- --- -## Procedure +## Reviewing a language -**New language** — read `en/common.json` *with* descriptions → settle your Tier C terms → write `i18n/locales//common.json` (same keys and order as `en`, `message` only, placeholders & brands preserved) → add a row to `_index.json` (`{"code","displayName"` = native name`,"englishName"}`) → run the QA list. Use the locale-code style the existing entries use (e.g. `fr`, `pt`, `zh-CN`). +**On Crowdin:** proofread in the editor — context, glossary highlights, and QA flags sit inline next to each string. -**Review (de / hu / …)** — read source and target side by side; for each key check glossary conformance (e.g. de `Exit-Node` → `Exit Node`, hu `Kilépő csomópont` → `Exit Node`), placeholder/`\n` integrity, consistency, tone, and that the meaning matches the English `description`. Fix in place, then report what you changed (especially term standardizations) so a native speaker can sanity-check. +**In the repo** — e.g. driving an AI agent with *"Read `i18n/TRANSLATING.md` and review the existing German translation"* — read source and target side by side; for each key check glossary conformance (e.g. de `Exit-Node` → `Exit Node`, hu `Kilépő csomópont` → `Exit Node`), placeholder/`\n` integrity, consistency, tone, and that the meaning matches the English `description`. Report what you found, and apply the fixes **on Crowdin** — direct edits to the locale files are overwritten by the next sync. --- ## QA before you finish -- [ ] Valid JSON · **every `en` key** present, same order · **no `description`** fields - [ ] Every `{placeholder}`, `\n`, and intentional space preserved · `...` / `… Failed` / `{name}` quotes kept -- [ ] Tier A/B left intact · Tier C applied consistently (and matching the existing bundle for your language) +- [ ] Tier A/B left intact · Tier C applied consistently (and matching the existing translation for your language) - [ ] Buttons & tray short · locale punctuation and capitalization applied -- [ ] New language added to `_index.json` +- [ ] Crowdin QA flags resolved (variables, glossary terms, punctuation) - [ ] **Tested in the running app** ↓ --- ## Test it in the app -A bundle can pass every check above and still read wrong on screen. **Run the app, switch to your language, and click through the real surfaces** — tray menu, main window, every Settings tab, the dialogs. Watch for text overflow or truncation, labels that are technically right but wrong *for what the control does*, leaked placeholders, and terms that drift between screens. +A translation can pass every check above and still read wrong on screen. **Run the app, switch to your language, and click through the real surfaces** — tray menu, main window, every Settings tab, the dialogs. Watch for text overflow or truncation, labels that are technically right but wrong *for what the control does*, leaked placeholders, and terms that drift between screens. How to run the app and switch language: see the project README. Can't run it (e.g. a headless agent)? Say so in your summary — don't silently skip this step. diff --git a/client/ui/i18n/check-translations.mjs b/client/ui/i18n/check-translations.mjs new file mode 100644 index 000000000..bd076e0e0 --- /dev/null +++ b/client/ui/i18n/check-translations.mjs @@ -0,0 +1,104 @@ +#!/usr/bin/env node +// Validates that every shipped translation bundle carries exactly the same set +// of keys as the English source of truth. English (en) defines the keys; every +// other locale declared in _index.json must match it 1:1: +// +// - no missing keys — a missing key silently falls back to English at runtime +// (see i18n bundle fallback), so the gap never surfaces to users or CI +// without this check; +// - no orphaned keys — keys left behind after an English key is renamed or +// removed are dead weight and a sign the locale is drifting. +// +// Pure Node, no dependencies, so it runs without installing the frontend +// toolchain. +// +// Local: node client/ui/i18n/check-translations.mjs (or: pnpm i18n:check) +// CI: .github/workflows/ui-translations.yml + +import { readdirSync, readFileSync } from "node:fs"; +import { dirname, join } from "node:path"; +import { fileURLToPath } from "node:url"; + +const SOURCE = "en"; +const localesDir = join(dirname(fileURLToPath(import.meta.url)), "locales"); +const isCI = Boolean(process.env.GITHUB_ACTIONS); + +function readJSON(path) { + return JSON.parse(readFileSync(path, "utf8")); +} + +function keysOf(langCode) { + return Object.keys(readJSON(join(localesDir, langCode, "common.json"))); +} + +// Emit a GitHub Actions annotation so failures render inline on the PR diff. +function annotate(file, message) { + if (isCI) console.log(`::error file=${file}::${message}`); +} + +const index = readJSON(join(localesDir, "_index.json")); +const declared = index.languages.map((l) => l.code); + +if (!declared.includes(SOURCE)) { + console.error(`FATAL: source language "${SOURCE}" is not declared in _index.json`); + process.exit(1); +} + +const sourceKeys = keysOf(SOURCE); +const sourceSet = new Set(sourceKeys); +console.log(`Source of truth: ${SOURCE}/common.json — ${sourceKeys.length} keys\n`); + +let failed = false; + +for (const code of declared) { + if (code === SOURCE) continue; + const file = `client/ui/i18n/locales/${code}/common.json`; + + let keys; + try { + keys = keysOf(code); + } catch (e) { + failed = true; + const msg = `bundle is declared in _index.json but common.json is missing or invalid (${e.message})`; + console.error(`✗ ${code}: ${msg}`); + annotate("client/ui/i18n/locales/_index.json", `${code}: ${msg}`); + continue; + } + + const set = new Set(keys); + const missing = sourceKeys.filter((k) => !set.has(k)); + const extra = keys.filter((k) => !sourceSet.has(k)); + + if (missing.length === 0 && extra.length === 0) { + console.log(`✓ ${code}: ${keys.length} keys`); + continue; + } + + failed = true; + console.error(`✗ ${code}: ${keys.length} keys (expected ${sourceKeys.length})`); + if (missing.length) { + console.error(` missing ${missing.length}: ${missing.join(", ")}`); + annotate(file, `Missing ${missing.length} key(s) present in ${SOURCE}: ${missing.join(", ")}`); + } + if (extra.length) { + console.error(` extra ${extra.length}: ${extra.join(", ")}`); + annotate(file, `Has ${extra.length} key(s) not present in ${SOURCE}: ${extra.join(", ")}`); + } +} + +// Locale directories present on disk but not declared in _index.json are never +// loaded by the app — surface them so dead translation files don't rot silently. +const onDisk = readdirSync(localesDir, { withFileTypes: true }) + .filter((e) => e.isDirectory()) + .map((e) => e.name); +const undeclared = onDisk.filter((d) => !declared.includes(d)); +if (undeclared.length) { + console.warn(`\n⚠ locale directories not declared in _index.json (not shipped): ${undeclared.join(", ")}`); +} + +console.log(); +if (failed) { + console.error("Translation check FAILED — every locale must match the English key set."); + process.exit(1); +} +console.log("Translation check passed — all locales match the English key set."); diff --git a/client/ui/i18n/locales/de/common.json b/client/ui/i18n/locales/de/common.json index 5e0d8096d..1208a37fe 100644 --- a/client/ui/i18n/locales/de/common.json +++ b/client/ui/i18n/locales/de/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Ändern des Autostarts fehlgeschlagen" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Nach dem Beenden verbunden bleiben", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "Die Verbindung bleibt im Hintergrund bestehen, nachdem Sie NetBird schließen. Sie endet erst, wenn Sie sie selbst trennen.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Anzeigesprache" }, @@ -756,7 +764,19 @@ "message": "Sensible Informationen anonymisieren" }, "settings.troubleshooting.anonymize.help": { - "message": "Versteckt öffentliche IP-Adressen und nicht-NetBird-Domains in Logs." + "message": "Verbirgt IP-Adressen, Domains und andere sensible Werte." + }, + "settings.troubleshooting.anonymize.info": { + "message": "Der Standardmodus lässt interne IPv4-Adressen und Peer-Namen für den Support lesbar. Der strikte Modus anonymisiert zusätzlich private (RFC 1918), CGNAT- und Link-Local-IP-Adressen, Peer-Namen und öffentliche WireGuard-Schlüssel. Wiederkehrende Werte erhalten denselben Platzhalter, sodass Peers unterscheidbar bleiben. Verwenden Sie den strikten Modus, wenn Sie das Debug-Paket außerhalb Ihrer Organisation weitergeben." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Keine" + }, + "settings.troubleshooting.anonymize.default": { + "message": "Standard" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Strikt" }, "settings.troubleshooting.systemInfo.label": { "message": "Systeminformationen einschließen" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Suchen Sie NetBird in der Taskleiste" }, + "welcome.titleMac": { + "message": "Suchen Sie NetBird in der Menüleiste" + }, "welcome.description": { "message": "NetBird läuft in Ihrer Taskleiste. Klicken Sie auf das Symbol, um sich zu verbinden, Profile zu wechseln oder die Einstellungen zu öffnen." }, + "welcome.descriptionMac": { + "message": "NetBird läuft in Ihrer Menüleiste. Klicken Sie auf das Symbol, um sich zu verbinden, Profile zu wechseln oder die Einstellungen zu öffnen." + }, "welcome.continue": { "message": "Weiter" }, @@ -1293,10 +1319,13 @@ "message": "Dokumentation" }, "daemon.outdated.title": { - "message": "NetBird-Dienst ist veraltet" + "message": "NetBird Client ist veraltet" }, "daemon.outdated.description": { - "message": "Aktualisieren Sie den NetBird-Dienst, um diese App zu verwenden." + "message": "Die neue GUI ist nicht mit Ihrem älteren Client kompatibel. Aktualisieren Sie Ihren Client, um die neue Anwendung zu verwenden." + }, + "daemon.outdated.download": { + "message": "Neueste Version herunterladen" }, "error.jwt_clock_skew": { "message": "Anmeldung fehlgeschlagen: Die Uhr dieses Geräts ist nicht mit dem Server synchron. Bitte synchronisieren Sie die Systemuhr und versuchen Sie es erneut." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "Vorgang fehlgeschlagen." + }, + "settings.ssh.privilege.hint": { + "message": "Erfordert {actor}. Führen Sie stattdessen dies aus:" + }, + "settings.ssh.privilege.oneWay": { + "message": "Sie können dies deaktivieren, aber zum erneuten Aktivieren sind {actor} erforderlich:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Sie können dies aktivieren, aber zum erneuten Deaktivieren sind {actor} erforderlich:" } } diff --git a/client/ui/i18n/locales/en/common.json b/client/ui/i18n/locales/en/common.json index 42d40ec30..694444497 100644 --- a/client/ui/i18n/locales/en/common.json +++ b/client/ui/i18n/locales/en/common.json @@ -735,6 +735,14 @@ "message": "Autostart Change Failed", "description": "Error-dialog title when changing the autostart setting fails." }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Stay Connected After Quitting", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "The connection stays up in the background after you close NetBird. It only stops when you disconnect it yourself.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Display Language", "description": "Label for the display-language picker." @@ -1005,11 +1013,27 @@ }, "settings.troubleshooting.anonymize.label": { "message": "Anonymize Sensitive Information", - "description": "Toggle label: anonymize sensitive information in the bundle." + "description": "Label for the anonymization level dropdown (None, Default, Strict)." }, "settings.troubleshooting.anonymize.help": { - "message": "Hides public IP addresses and non-NetBird domains from logs.", - "description": "Helper text for anonymizing logs (hides public IPs and non-NetBird domains)." + "message": "Hides IP addresses, domains, and other sensitive values.", + "description": "Helper text under the anonymization dropdown. The level details live in the info tooltip." + }, + "settings.troubleshooting.anonymize.info": { + "message": "Default keeps internal IPv4 addresses and peer names readable for support. Strict additionally anonymizes private (RFC 1918), CGNAT, and link-local IP addresses, peer names, and WireGuard public keys. Recurring values map to the same placeholder, so peers stay distinguishable. Use Strict when sharing the bundle outside your organization.", + "description": "Info tooltip explaining the anonymization levels. 'RFC 1918', 'CGNAT', 'link-local', and 'WireGuard' are technical terms — keep them." + }, + "settings.troubleshooting.anonymize.none": { + "message": "None", + "description": "Dropdown option: no anonymization." + }, + "settings.troubleshooting.anonymize.default": { + "message": "Default", + "description": "Dropdown option: default anonymization level." + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Strict", + "description": "Dropdown option: strict anonymization level." }, "settings.troubleshooting.systemInfo.label": { "message": "Include System Information", @@ -1377,11 +1401,19 @@ }, "welcome.title": { "message": "Look for NetBird in your tray", - "description": "Heading on the first onboarding step, pointing the user to the tray icon. 'tray' = system tray / menu bar." + "description": "Heading on the first onboarding step, pointing the user to the tray icon. Shown on Windows and Linux; macOS uses welcome.titleMac." + }, + "welcome.titleMac": { + "message": "Look for NetBird in your menu bar", + "description": "Heading on the first onboarding step on macOS, pointing the user to the menu bar icon. Use your language's Apple term for the macOS menu bar." }, "welcome.description": { "message": "NetBird lives in your tray. Click the icon to connect, switch profiles, or open settings.", - "description": "Body of the first onboarding step explaining the tray icon." + "description": "Body of the first onboarding step explaining the tray icon. Shown on Windows and Linux; macOS uses welcome.descriptionMac." + }, + "welcome.descriptionMac": { + "message": "NetBird lives in your menu bar. Click the icon to connect, switch profiles, or open settings.", + "description": "Body of the first onboarding step on macOS explaining the menu bar icon. Use your language's Apple term for the macOS menu bar." }, "welcome.continue": { "message": "Continue", @@ -1724,12 +1756,16 @@ "description": "Documentation link on the daemon-unavailable overlay." }, "daemon.outdated.title": { - "message": "NetBird Service Is Outdated", - "description": "Title of the overlay shown when the NetBird background service is too old to drive this UI." + "message": "NetBird Client Is Outdated", + "description": "Title of the overlay shown when the NetBird client (daemon) is too old to drive this UI." }, "daemon.outdated.description": { - "message": "Update the NetBird service to use this app.", - "description": "Body of the daemon-outdated overlay telling the user to upgrade the service." + "message": "The new GUI isn't compatible with the older NetBird client. Update your client to use the new application.", + "description": "Body of the daemon-outdated overlay explaining that the GUI is newer than the client and the client must be updated." + }, + "daemon.outdated.download": { + "message": "Download Latest", + "description": "Button on the daemon-outdated overlay that opens the download page for the latest release." }, "error.jwt_clock_skew": { "message": "Sign-in failed: this device's clock is out of sync with the server. Please sync your system clock and try again.", @@ -1762,5 +1798,17 @@ "error.unknown": { "message": "Operation failed.", "description": "Generic fallback error message used when no specific error applies." + }, + "settings.ssh.privilege.hint": { + "message": "Requires {actor}. Run this instead:", + "description": "Help text under an SSH setting the user cannot change: it needs elevated privileges. {actor} is 'root' on Linux/macOS or 'administrator privileges' on Windows. Followed by a copyable command." + }, + "settings.ssh.privilege.oneWay": { + "message": "You can switch this off, but switching it back on needs {actor}:", + "description": "Warning under an SSH setting an unprivileged user may disable but not re-enable. {actor} is 'root' on Linux/macOS or 'administrator privileges' on Windows. Followed by a copyable command." + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "You can switch this on, but switching it back off needs {actor}:", + "description": "Warning under the SSH authentication setting, which an unprivileged user may re-enable but not disable again. {actor} is 'root' on Linux/macOS or 'administrator privileges' on Windows. Followed by a copyable command." } } diff --git a/client/ui/i18n/locales/es/common.json b/client/ui/i18n/locales/es/common.json index 47faee61f..6dc4ffd0b 100644 --- a/client/ui/i18n/locales/es/common.json +++ b/client/ui/i18n/locales/es/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Error al cambiar el inicio automático" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Permanecer conectado al salir", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "La conexión sigue activa en segundo plano después de cerrar NetBird. Solo se detiene cuando la desconectas tú.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Idioma de la interfaz" }, @@ -756,7 +764,19 @@ "message": "Anonimizar información sensible" }, "settings.troubleshooting.anonymize.help": { - "message": "Oculta las direcciones IP públicas y los dominios ajenos a NetBird de los registros." + "message": "Oculta direcciones IP, dominios y otros valores sensibles." + }, + "settings.troubleshooting.anonymize.info": { + "message": "El modo predeterminado mantiene legibles las direcciones IPv4 internas y los nombres de los peers para el soporte. El modo estricto anonimiza además las direcciones IP privadas (RFC 1918), CGNAT y de enlace local, los nombres de los peers y las claves públicas de WireGuard. Los valores recurrentes se asignan al mismo marcador de posición, por lo que los peers siguen siendo distinguibles. Use el modo estricto cuando comparta el paquete de diagnóstico fuera de su organización." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Ninguno" + }, + "settings.troubleshooting.anonymize.default": { + "message": "Predeterminado" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Estricto" }, "settings.troubleshooting.systemInfo.label": { "message": "Incluir información del sistema" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Busque NetBird en su bandeja del sistema" }, + "welcome.titleMac": { + "message": "Busque NetBird en su barra de menús" + }, "welcome.description": { "message": "NetBird reside en su bandeja del sistema. Haga clic en el icono para conectarse, cambiar de perfil o abrir la configuración." }, + "welcome.descriptionMac": { + "message": "NetBird reside en su barra de menús. Haga clic en el icono para conectarse, cambiar de perfil o abrir la configuración." + }, "welcome.continue": { "message": "Continuar" }, @@ -1293,10 +1319,13 @@ "message": "Documentación" }, "daemon.outdated.title": { - "message": "El servicio de NetBird está desactualizado" + "message": "NetBird Client está desactualizado" }, "daemon.outdated.description": { - "message": "Actualice el servicio de NetBird para usar esta aplicación." + "message": "La nueva GUI no es compatible con su cliente anterior. Actualice su cliente para usar la nueva aplicación." + }, + "daemon.outdated.download": { + "message": "Descargar la última versión" }, "error.jwt_clock_skew": { "message": "Error al iniciar sesión: el reloj de este dispositivo no está sincronizado con el servidor. Sincronice el reloj del sistema e inténtelo de nuevo." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "La operación falló." + }, + "settings.ssh.privilege.hint": { + "message": "Requiere {actor}. Ejecute esto en su lugar:" + }, + "settings.ssh.privilege.oneWay": { + "message": "Puede desactivarlo, pero volver a activarlo requiere {actor}:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Puede activarlo, pero volver a desactivarlo requiere {actor}:" } } diff --git a/client/ui/i18n/locales/fr/common.json b/client/ui/i18n/locales/fr/common.json index be0836e93..d3e54440c 100644 --- a/client/ui/i18n/locales/fr/common.json +++ b/client/ui/i18n/locales/fr/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Échec de la modification du démarrage automatique" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Rester connecté après la fermeture", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "La connexion reste active en arrière-plan après la fermeture de NetBird. Elle ne s'arrête que si vous la coupez vous-même.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Langue d’affichage" }, @@ -756,7 +764,19 @@ "message": "Anonymiser les informations sensibles" }, "settings.troubleshooting.anonymize.help": { - "message": "Masque les adresses IP publiques et les domaines non-NetBird dans les journaux." + "message": "Masque les adresses IP, les domaines et d'autres valeurs sensibles." + }, + "settings.troubleshooting.anonymize.info": { + "message": "Le mode par défaut garde les adresses IPv4 internes et les noms des pairs lisibles pour le support. Le mode strict anonymise en plus les adresses IP privées (RFC 1918), CGNAT et de lien local, les noms des pairs et les clés publiques WireGuard. Les valeurs récurrentes reçoivent le même espace réservé, les pairs restent donc distinguables. Utilisez le mode strict lorsque vous partagez le lot de diagnostic en dehors de votre organisation." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Aucune" + }, + "settings.troubleshooting.anonymize.default": { + "message": "Par défaut" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Strict" }, "settings.troubleshooting.systemInfo.label": { "message": "Inclure les informations système" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Cherchez NetBird dans votre barre d’état système" }, + "welcome.titleMac": { + "message": "Cherchez NetBird dans votre barre des menus" + }, "welcome.description": { "message": "NetBird se trouve dans votre barre d’état système. Cliquez sur l’icône pour vous connecter, changer de profil ou ouvrir les paramètres." }, + "welcome.descriptionMac": { + "message": "NetBird se trouve dans votre barre des menus. Cliquez sur l’icône pour vous connecter, changer de profil ou ouvrir les paramètres." + }, "welcome.continue": { "message": "Continuer" }, @@ -1293,10 +1319,13 @@ "message": "Documentation" }, "daemon.outdated.title": { - "message": "Le service NetBird est obsolète" + "message": "Le Client NetBird est obsolète" }, "daemon.outdated.description": { - "message": "Mettez à jour le service NetBird pour utiliser cette application." + "message": "La nouvelle GUI n'est pas compatible avec votre ancien client. Mettez à jour votre client pour utiliser la nouvelle application." + }, + "daemon.outdated.download": { + "message": "Télécharger la dernière version" }, "error.jwt_clock_skew": { "message": "Échec de la connexion : l’horloge de cet appareil n’est pas synchronisée avec le serveur. Veuillez synchroniser l’horloge de votre système et réessayer." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "L’opération a échoué." + }, + "settings.ssh.privilege.hint": { + "message": "Nécessite {actor}. Exécutez plutôt ceci :" + }, + "settings.ssh.privilege.oneWay": { + "message": "Vous pouvez le désactiver, mais le réactiver nécessite {actor} :" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Vous pouvez l’activer, mais le désactiver de nouveau nécessite {actor} :" } } diff --git a/client/ui/i18n/locales/hu/common.json b/client/ui/i18n/locales/hu/common.json index b54918364..19aede17f 100644 --- a/client/ui/i18n/locales/hu/common.json +++ b/client/ui/i18n/locales/hu/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Az automatikus indítás módosítása sikertelen" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Kapcsolat megtartása kilépéskor", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "A kapcsolat a háttérben megmarad, miután bezárod a NetBirdöt. Csak akkor szakad meg, ha te magad bontod.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Megjelenítési nyelv" }, @@ -756,7 +764,19 @@ "message": "Érzékeny információk anonimizálása" }, "settings.troubleshooting.anonymize.help": { - "message": "Elrejti a nyilvános IP-címeket és a nem-NetBird tartományokat a naplókban." + "message": "Elrejti az IP-címeket, a tartományokat és más érzékeny értékeket." + }, + "settings.troubleshooting.anonymize.info": { + "message": "Az Alapértelmezett szint a belső IPv4-címeket és a peer-neveket olvashatóan hagyja a támogatás számára. A Szigorú ezen felül anonimizálja a privát (RFC 1918), CGNAT és link-local IP-címeket, a peer-neveket és a WireGuard nyilvános kulcsokat. Az ismétlődő értékek ugyanazt a helyettesítőt kapják, így a peerek megkülönböztethetők maradnak. Használja a Szigorú szintet, ha a hibakeresési csomagot a szervezetén kívül osztja meg." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Nincs" + }, + "settings.troubleshooting.anonymize.default": { + "message": "Alapértelmezett" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Szigorú" }, "settings.troubleshooting.systemInfo.label": { "message": "Rendszerinformációk beillesztése" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Keresse a NetBirdöt a tálcán" }, + "welcome.titleMac": { + "message": "Keresse a NetBirdöt a menüsorban" + }, "welcome.description": { "message": "A NetBird a tálcán fut. Kattintson az ikonra a csatlakozáshoz, profilváltáshoz vagy a beállítások megnyitásához." }, + "welcome.descriptionMac": { + "message": "A NetBird a menüsorban fut. Kattintson az ikonra a csatlakozáshoz, profilváltáshoz vagy a beállítások megnyitásához." + }, "welcome.continue": { "message": "Folytatás" }, @@ -1293,10 +1319,13 @@ "message": "Dokumentáció" }, "daemon.outdated.title": { - "message": "A NetBird szolgáltatás elavult" + "message": "A NetBird Kliens elavult" }, "daemon.outdated.description": { - "message": "Frissítsd a NetBird szolgáltatást az alkalmazás használatához." + "message": "Az új GUI nem kompatibilis a régebbi klienseddel. Frissítsd a klienst az új alkalmazás használatához." + }, + "daemon.outdated.download": { + "message": "Legújabb letöltése" }, "error.jwt_clock_skew": { "message": "A bejelentkezés sikertelen: az eszköz órája eltér a szerverétől. Kérjük, szinkronizálja a rendszer óráját, majd próbálja újra." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "A művelet meghiúsult." + }, + "settings.ssh.privilege.hint": { + "message": "{actor} szükséges hozzá. Futtassa inkább ezt:" + }, + "settings.ssh.privilege.oneWay": { + "message": "Kikapcsolhatja, de a visszakapcsolásához {actor} szükséges:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Bekapcsolhatja, de az ismételt kikapcsolásához {actor} szükséges:" } } diff --git a/client/ui/i18n/locales/it/common.json b/client/ui/i18n/locales/it/common.json index 603364fa2..dab9e0cb4 100644 --- a/client/ui/i18n/locales/it/common.json +++ b/client/ui/i18n/locales/it/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Modifica avvio automatico non riuscita" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Resta connesso dopo la chiusura", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "La connessione resta attiva in background dopo la chiusura di NetBird. Si interrompe solo quando la disconnetti tu.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Lingua dell'interfaccia" }, @@ -756,7 +764,19 @@ "message": "Anonimizza informazioni sensibili" }, "settings.troubleshooting.anonymize.help": { - "message": "Nasconde gli indirizzi IP pubblici e i domini non NetBird dai log." + "message": "Nasconde indirizzi IP, domini e altri valori sensibili." + }, + "settings.troubleshooting.anonymize.info": { + "message": "La modalità predefinita mantiene leggibili gli indirizzi IPv4 interni e i nomi dei peer per il supporto. La modalità rigorosa anonimizza inoltre gli indirizzi IP privati (RFC 1918), CGNAT e link-local, i nomi dei peer e le chiavi pubbliche WireGuard. I valori ricorrenti vengono associati allo stesso segnaposto, quindi i peer restano distinguibili. Usa la modalità rigorosa quando condividi il pacchetto di debug al di fuori della tua organizzazione." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Nessuna" + }, + "settings.troubleshooting.anonymize.default": { + "message": "Predefinito" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Rigoroso" }, "settings.troubleshooting.systemInfo.label": { "message": "Includi informazioni di sistema" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Cerchi NetBird nella tray" }, + "welcome.titleMac": { + "message": "Cerchi NetBird nella barra dei menu" + }, "welcome.description": { "message": "NetBird risiede nella tray. Clicchi sull'icona per connettersi, cambiare profilo o aprire le impostazioni." }, + "welcome.descriptionMac": { + "message": "NetBird risiede nella barra dei menu. Clicchi sull'icona per connettersi, cambiare profilo o aprire le impostazioni." + }, "welcome.continue": { "message": "Continua" }, @@ -1293,10 +1319,13 @@ "message": "Documentazione" }, "daemon.outdated.title": { - "message": "Il servizio NetBird è obsoleto" + "message": "NetBird Client è obsoleto" }, "daemon.outdated.description": { - "message": "Aggiorna il servizio NetBird per usare questa app." + "message": "La nuova GUI non è compatibile con il tuo client precedente. Aggiorna il client per usare la nuova applicazione." + }, + "daemon.outdated.download": { + "message": "Scarica l'ultima versione" }, "error.jwt_clock_skew": { "message": "Accesso non riuscito: l'orologio di questo dispositivo non è sincronizzato con il server. Sincronizzi l'orologio di sistema e riprovi." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "Operazione non riuscita." + }, + "settings.ssh.privilege.hint": { + "message": "Richiede {actor}. Esegua invece questo:" + }, + "settings.ssh.privilege.oneWay": { + "message": "Può disabilitarlo, ma riabilitarlo richiede {actor}:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Può abilitarlo, ma disabilitarlo di nuovo richiede {actor}:" } } diff --git a/client/ui/i18n/locales/ja/common.json b/client/ui/i18n/locales/ja/common.json index cd54bce17..246c232a8 100644 --- a/client/ui/i18n/locales/ja/common.json +++ b/client/ui/i18n/locales/ja/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "自動起動の変更に失敗しました" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "終了後も接続を維持", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "NetBird を閉じたあとも接続はバックグラウンドで維持されます。自分で切断したときにだけ停止します。", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "表示言語" }, @@ -756,7 +764,19 @@ "message": "機密情報を匿名化" }, "settings.troubleshooting.anonymize.help": { - "message": "ログからパブリック IP アドレスと NetBird 以外のドメインを隠します。" + "message": "IP アドレス、ドメイン、その他の機密性の高い値を隠します。" + }, + "settings.troubleshooting.anonymize.info": { + "message": "「デフォルト」では、サポートのために内部 IPv4 アドレスとピア名は読める状態のまま残ります。「厳格」では、さらにプライベート (RFC 1918)、CGNAT、リンクローカルの IP アドレス、ピア名、WireGuard 公開鍵も匿名化されます。繰り返し現れる値は同じプレースホルダーに置き換えられるため、ピアは区別できます。デバッグバンドルを組織外に共有する場合は「厳格」を使用してください。" + }, + "settings.troubleshooting.anonymize.none": { + "message": "なし" + }, + "settings.troubleshooting.anonymize.default": { + "message": "デフォルト" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "厳格" }, "settings.troubleshooting.systemInfo.label": { "message": "システム情報を含める" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "トレイの NetBird を確認してください" }, + "welcome.titleMac": { + "message": "メニューバーの NetBird を確認してください" + }, "welcome.description": { "message": "NetBird はトレイに常駐します。アイコンをクリックして、接続、プロファイルの切り替え、設定を開くことができます。" }, + "welcome.descriptionMac": { + "message": "NetBird はメニューバーに常駐します。アイコンをクリックして、接続、プロファイルの切り替え、設定を開くことができます。" + }, "welcome.continue": { "message": "続ける" }, @@ -1298,6 +1324,9 @@ "daemon.outdated.description": { "message": "このアプリを使用するには NetBird サービスを更新してください。" }, + "daemon.outdated.download": { + "message": "最新版をダウンロード" + }, "error.jwt_clock_skew": { "message": "サインインに失敗しました: このデバイスの時計がサーバーと同期していません。システムの時計を同期してからもう一度お試しください。" }, @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "操作に失敗しました。" + }, + "settings.ssh.privilege.hint": { + "message": "{actor}が必要です。代わりに次のコマンドを実行してください:" + }, + "settings.ssh.privilege.oneWay": { + "message": "無効にはできますが、再度有効にするには{actor}が必要です:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "有効にはできますが、再度無効にするには{actor}が必要です:" } } diff --git a/client/ui/i18n/locales/pt/common.json b/client/ui/i18n/locales/pt/common.json index 2ed0a94c5..418e93717 100644 --- a/client/ui/i18n/locales/pt/common.json +++ b/client/ui/i18n/locales/pt/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Falha ao alterar o início automático" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Permanecer conectado ao sair", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "A conexão continua ativa em segundo plano depois de fechar o NetBird. Ela só para quando você mesmo a desconecta.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Idioma de exibição" }, @@ -756,7 +764,19 @@ "message": "Anonimizar informações sensíveis" }, "settings.troubleshooting.anonymize.help": { - "message": "Oculta endereços IP públicos e domínios que não são do NetBird nos logs." + "message": "Oculta endereços IP, domínios e outros valores sensíveis." + }, + "settings.troubleshooting.anonymize.info": { + "message": "O modo padrão mantém os endereços IPv4 internos e os nomes dos peers legíveis para o suporte. O modo estrito anonimiza também os endereços IP privados (RFC 1918), CGNAT e link-local, os nomes dos peers e as chaves públicas do WireGuard. Valores recorrentes recebem o mesmo marcador, então os peers continuam distinguíveis. Use o modo estrito ao compartilhar o pacote de depuração fora da sua organização." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Nenhum" + }, + "settings.troubleshooting.anonymize.default": { + "message": "Padrão" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Estrito" }, "settings.troubleshooting.systemInfo.label": { "message": "Incluir informações do sistema" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Procure o NetBird na sua bandeja" }, + "welcome.titleMac": { + "message": "Procure o NetBird na sua barra de menus" + }, "welcome.description": { "message": "O NetBird fica na sua bandeja. Clique no ícone para conectar, alternar perfis ou abrir as configurações." }, + "welcome.descriptionMac": { + "message": "O NetBird fica na sua barra de menus. Clique no ícone para conectar, alternar perfis ou abrir as configurações." + }, "welcome.continue": { "message": "Continuar" }, @@ -1293,10 +1319,13 @@ "message": "Documentação" }, "daemon.outdated.title": { - "message": "O serviço NetBird está desatualizado" + "message": "O NetBird Client está desatualizado" }, "daemon.outdated.description": { - "message": "Atualize o serviço NetBird para usar este aplicativo." + "message": "A nova GUI não é compatível com o seu cliente mais antigo. Atualize o seu cliente para usar o novo aplicativo." + }, + "daemon.outdated.download": { + "message": "Baixar a versão mais recente" }, "error.jwt_clock_skew": { "message": "Falha no login: o relógio deste dispositivo está fora de sincronia com o servidor. Sincronize o relógio do sistema e tente novamente." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "A operação falhou." + }, + "settings.ssh.privilege.hint": { + "message": "Requer {actor}. Execute isto em vez disso:" + }, + "settings.ssh.privilege.oneWay": { + "message": "Você pode desativar isto, mas ativar novamente requer {actor}:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Você pode ativar isto, mas desativar novamente requer {actor}:" } } diff --git a/client/ui/i18n/locales/ru/common.json b/client/ui/i18n/locales/ru/common.json index 6ba7de8cc..958b5a21c 100644 --- a/client/ui/i18n/locales/ru/common.json +++ b/client/ui/i18n/locales/ru/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "Не удалось изменить автозапуск" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "Оставаться подключённым после выхода", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "Соединение остаётся активным в фоне после закрытия NetBird. Оно прервётся, только когда вы отключите его сами.", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "Язык интерфейса" }, @@ -756,7 +764,19 @@ "message": "Анонимизировать конфиденциальную информацию" }, "settings.troubleshooting.anonymize.help": { - "message": "Скрывает публичные IP-адреса и сторонние (не относящиеся к NetBird) домены в журналах." + "message": "Скрывает IP-адреса, домены и другие конфиденциальные значения." + }, + "settings.troubleshooting.anonymize.info": { + "message": "Режим «По умолчанию» оставляет внутренние IPv4-адреса и имена пиров читаемыми для поддержки. Режим «Строгий» дополнительно анонимизирует частные (RFC 1918), CGNAT и link-local IP-адреса, имена пиров и публичные ключи WireGuard. Повторяющиеся значения заменяются одним и тем же заполнителем, поэтому пиры остаются различимыми. Используйте режим «Строгий», когда передаёте отладочный пакет за пределы вашей организации." + }, + "settings.troubleshooting.anonymize.none": { + "message": "Нет" + }, + "settings.troubleshooting.anonymize.default": { + "message": "По умолчанию" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "Строгий" }, "settings.troubleshooting.systemInfo.label": { "message": "Включить сведения о системе" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "Найдите NetBird в системном трее" }, + "welcome.titleMac": { + "message": "Найдите NetBird в строке меню" + }, "welcome.description": { "message": "NetBird находится в системном трее. Нажмите на значок, чтобы подключиться, переключить профиль или открыть настройки." }, + "welcome.descriptionMac": { + "message": "NetBird находится в строке меню. Нажмите на значок, чтобы подключиться, переключить профиль или открыть настройки." + }, "welcome.continue": { "message": "Продолжить" }, @@ -1293,10 +1319,13 @@ "message": "Документация" }, "daemon.outdated.title": { - "message": "Служба NetBird устарела" + "message": "Клиент NetBird устарел" }, "daemon.outdated.description": { - "message": "Обновите службу NetBird, чтобы использовать это приложение." + "message": "Новый GUI несовместим с вашим более старым клиентом. Обновите клиент, чтобы использовать новое приложение." + }, + "daemon.outdated.download": { + "message": "Скачать последнюю версию" }, "error.jwt_clock_skew": { "message": "Не удалось войти: часы этого устройства рассинхронизированы с сервером. Синхронизируйте системные часы и повторите попытку." @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "Не удалось выполнить операцию." + }, + "settings.ssh.privilege.hint": { + "message": "Требуются {actor}. Выполните вместо этого:" + }, + "settings.ssh.privilege.oneWay": { + "message": "Отключить можно, но чтобы включить снова, нужны {actor}:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "Включить можно, но чтобы отключить снова, нужны {actor}:" } } diff --git a/client/ui/i18n/locales/zh-CN/common.json b/client/ui/i18n/locales/zh-CN/common.json index 609344fc0..90ae5e003 100644 --- a/client/ui/i18n/locales/zh-CN/common.json +++ b/client/ui/i18n/locales/zh-CN/common.json @@ -551,6 +551,14 @@ "settings.general.autostart.errorTitle": { "message": "更改自启动设置失败" }, + "settings.general.keepConnectedOnQuit.label": { + "message": "退出后保持连接", + "description": "Toggle label: keep the VPN connection up after quitting the UI." + }, + "settings.general.keepConnectedOnQuit.help": { + "message": "关闭 NetBird 后,连接会在后台保持。只有你自己断开时才会停止。", + "description": "Helper text for the stay-connected-after-quitting toggle." + }, "settings.general.language.label": { "message": "显示语言" }, @@ -756,7 +764,19 @@ "message": "匿名化敏感信息" }, "settings.troubleshooting.anonymize.help": { - "message": "从日志中隐藏公共 IP 地址和非 NetBird 域名。" + "message": "隐藏 IP 地址、域名和其他敏感值。" + }, + "settings.troubleshooting.anonymize.info": { + "message": "默认级别保留内部 IPv4 地址和对等节点名称,便于支持人员阅读。严格级别还会匿名化私有 (RFC 1918)、CGNAT 和链路本地 IP 地址、对等节点名称以及 WireGuard 公钥。相同的值会映射到相同的占位符,因此对等节点仍可区分。向组织外部分享调试包时请使用严格级别。" + }, + "settings.troubleshooting.anonymize.none": { + "message": "无" + }, + "settings.troubleshooting.anonymize.default": { + "message": "默认" + }, + "settings.troubleshooting.anonymize.strict": { + "message": "严格" }, "settings.troubleshooting.systemInfo.label": { "message": "包含系统信息" @@ -1034,9 +1054,15 @@ "welcome.title": { "message": "在托盘中查找 NetBird" }, + "welcome.titleMac": { + "message": "在菜单栏中查找 NetBird" + }, "welcome.description": { "message": "NetBird 驻留在您的托盘中。点击图标即可连接、切换配置文件或打开设置。" }, + "welcome.descriptionMac": { + "message": "NetBird 驻留在您的菜单栏中。点击图标即可连接、切换配置文件或打开设置。" + }, "welcome.continue": { "message": "继续" }, @@ -1293,10 +1319,13 @@ "message": "文档" }, "daemon.outdated.title": { - "message": "NetBird 服务版本过旧" + "message": "NetBird 客户端版本过旧" }, "daemon.outdated.description": { - "message": "请更新 NetBird 服务以使用此应用。" + "message": "新版 GUI 与您较旧的客户端不兼容。请更新客户端以使用新应用。" + }, + "daemon.outdated.download": { + "message": "下载最新版本" }, "error.jwt_clock_skew": { "message": "登录失败:此设备的时钟与服务器不同步。请同步您的系统时钟后重试。" @@ -1321,5 +1350,14 @@ }, "error.unknown": { "message": "操作失败。" + }, + "settings.ssh.privilege.hint": { + "message": "需要{actor}。请改为运行:" + }, + "settings.ssh.privilege.oneWay": { + "message": "您可以关闭此项,但重新开启需要{actor}:" + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "您可以开启此项,但再次关闭需要{actor}:" } } diff --git a/client/ui/main.go b/client/ui/main.go index 4889bad79..e20bfe074 100644 --- a/client/ui/main.go +++ b/client/ui/main.go @@ -14,7 +14,6 @@ import ( "github.com/sirupsen/logrus" "github.com/wailsapp/wails/v3/pkg/application" "github.com/wailsapp/wails/v3/pkg/events" - "github.com/wailsapp/wails/v3/pkg/services/notifications" "github.com/netbirdio/netbird/client/ui/authsession" "github.com/netbirdio/netbird/client/ui/i18n" @@ -63,7 +62,7 @@ type registeredServices struct { profiles *services.Profiles update *services.Update daemonFeed *services.DaemonFeed - notifier *notifications.NotificationService + notifier *Notifier compat *services.Compat profileSwitcher *services.ProfileSwitcher bundle *i18n.Bundle @@ -96,14 +95,13 @@ func main() { } }) - settings := services.NewSettings(conn) profiles := services.NewProfiles(conn) // updater.Holder owns the typed update State; DaemonFeed feeds it and the // Update service is a thin Wails-bound facade over it plus the install RPCs. updaterHolder := updater.NewHolder(app.Event) update := services.NewUpdate(conn, updaterHolder) daemonFeed := services.NewDaemonFeed(conn, app.Event, updaterHolder, debugLog) - notifier := notifications.New() + notifier := newNotifier() compat := services.NewCompat(conn) // macOS shows no toast until permission is requested. Run it after // ApplicationStarted so the notifier's Startup has initialised the @@ -117,6 +115,7 @@ func main() { bundle, prefStore, localizer := buildI18n(app) // After bundle + prefStore: both are used to localise daemon errors. + settings := services.NewSettings(conn, bundle, prefStore, daemonAddr) connection := services.NewConnection(conn, bundle, prefStore) profileSwitcher := services.NewProfileSwitcher(profiles, connection, daemonFeed) // authsession.Session owns the full extend + dismiss surface the tray @@ -140,13 +139,11 @@ func main() { prefStore: prefStore, }) - window := newMainWindow(app, prefStore) - - // Settings is created eagerly (hidden) so the first gear click paints - // instantly and React keeps per-tab state across reopens. The other - // auxiliary windows stay lazy + destroy-on-close so Wails's macOS - // dock-reopen handler can't resurrect them. - windowManager := services.NewWindowManager(app, window, bundle, prefStore, iconWindow) + windowManager := services.NewWindowManager(app, nil, bundle, prefStore, iconWindow) + windowManager.SetMainFactory(func(startURL string) *application.WebviewWindow { + return newMainWindow(app, prefStore, windowManager, startURL) + }) + registerDockReopenHook(app, windowManager) // Minimal WMs (XEmbed-tray path) neither center small windows nor restore // position across hide -> show, dropping them top-left. Gate Go-side // re-centering on that environment; nil leaves placement to the WM on full @@ -169,7 +166,7 @@ func main() { // RegisterStatusNotifierItem hits a watcher we control. startStatusNotifierWatcher() - tray = NewTray(app, window, TrayServices{ + tray = NewTray(app, nil, TrayServices{ Connection: connection, Settings: settings, Profiles: profiles, @@ -181,6 +178,7 @@ func main() { WindowManager: windowManager, Session: authSession, Localizer: localizer, + Preferences: prefStore, }) listenForShowSignal(context.Background(), tray) @@ -210,7 +208,7 @@ func main() { // requestNotificationAuthorization prompts for macOS notification permission. // The request blocks until the user responds (up to 3 minutes), so callers run // it in a goroutine. No-op on Linux/Windows. -func requestNotificationAuthorization(notifier *notifications.NotificationService) { +func requestNotificationAuthorization(notifier *Notifier) { authorized, err := notifier.CheckNotificationAuthorization() if err != nil { logrus.Debugf("check notification authorization: %v", err) @@ -279,7 +277,12 @@ func newApplication(onSecondInstance func()) *application.App { ActivationPolicy: application.ActivationPolicyAccessory, }, Linux: application.LinuxOptions{ - ProgramName: "netbird", + ProgramName: "netbird", + DisableQuitOnLastWindowClosed: true, + }, + Windows: application.WindowsOptions{ + WndProcInterceptor: endSessionInterceptor(), + DisableQuitOnLastWindowClosed: true, }, SingleInstance: &application.SingleInstanceOptions{ UniqueID: "io.netbird.ui", @@ -335,9 +338,7 @@ func registerServices(app *application.App, conn *Conn, s registeredServices) { app.RegisterService(application.NewService(s.compat)) } -// newMainWindow creates the hidden main window, sized to the user's last view -// mode, and installs the hide-on-close and macOS dock-reopen hooks. -func newMainWindow(app *application.App, prefStore *preferences.Store) *application.WebviewWindow { +func newMainWindow(app *application.App, prefStore *preferences.Store, wm *services.WindowManager, startURL string) *application.WebviewWindow { // Width matches the last view mode so Advanced-mode users don't see the // window pop from 380px to 900px on launch. Height is mode-agnostic. initialWidth := 380 @@ -354,7 +355,7 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat InitialPosition: application.WindowCentered, Hidden: true, BackgroundColour: services.WindowBackgroundColour, - URL: "/", + URL: startURL, DisableResize: true, MinimiseButtonState: application.ButtonHidden, MaximiseButtonState: application.ButtonHidden, @@ -365,26 +366,25 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat }, }) - // Hide instead of quit on close; "really quit" is reached via tray -> Quit. - window.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) { - e.Cancel() - window.Hide() + window.RegisterHook(events.Common.WindowClosing, func(_ *application.WindowEvent) { + if services.ShuttingDown() { + return + } + wm.ForgetMain() }) - // On macOS, Wails' default applicationShouldHandleReopen handler Show()s - // every hidden window on dock-icon click, resurrecting hide-on-close - // surfaces like Settings. Cancel it in a hook (hooks run before listeners) - // and show only the main window. No-op elsewhere — the event never fires. - if runtime.GOOS == "darwin" { - app.Event.RegisterApplicationEventHook(events.Mac.ApplicationShouldHandleReopen, func(e *application.ApplicationEvent) { - e.Cancel() - if e.Context().HasVisibleWindows() { - return - } - window.Show() - window.Focus() - }) - } - return window } + +func registerDockReopenHook(app *application.App, wm *services.WindowManager) { + if runtime.GOOS != "darwin" { + return + } + app.Event.RegisterApplicationEventHook(events.Mac.ApplicationShouldHandleReopen, func(e *application.ApplicationEvent) { + if e.Context().HasVisibleWindows() { + return + } + e.Cancel() + wm.ShowMain() + }) +} diff --git a/client/ui/netbird-ui.rb.tmpl b/client/ui/netbird-ui.rb.tmpl index 06971909d..1c77e6717 100644 --- a/client/ui/netbird-ui.rb.tmpl +++ b/client/ui/netbird-ui.rb.tmpl @@ -29,8 +29,13 @@ cask "{{ $projectName }}" do end uninstall_preflight do - system_command "#{appdir}/Netbird UI.app/uninstaller.sh", - sudo: false + system_command "/bin/sh", + args: ["-c", <<~CMD], + launchctl bootout system/netbird 2>/dev/null || \ + launchctl unload /Library/LaunchDaemons/netbird.plist 2>/dev/null || true + rm -f /Library/LaunchDaemons/netbird.plist + CMD + sudo: true end name "Netbird UI" diff --git a/client/ui/notifier.go b/client/ui/notifier.go new file mode 100644 index 000000000..71ae3b0df --- /dev/null +++ b/client/ui/notifier.go @@ -0,0 +1,101 @@ +//go:build !android && !ios && !freebsd && !js + +package main + +import ( + "context" + "errors" + "sync/atomic" + + log "github.com/sirupsen/logrus" + "github.com/wailsapp/wails/v3/pkg/application" + "github.com/wailsapp/wails/v3/pkg/services/notifications" +) + +var errNotificationsUnavailable = errors.New("notifications unavailable") + +// Notifier wraps the Wails notification service so an unavailable backend +// disables notifications instead of aborting the app. Startup fails for +// environment reasons (a bare unbundled binary on macOS has no bundle +// identifier, a headless Linux session has no D-Bus session bus), and Wails +// treats a service startup error as fatal. After a failed startup every call +// is a no-op: on macOS, touching UNUserNotificationCenter without a bundle +// identifier raises an Objective-C exception that recover() cannot catch. +type Notifier struct { + inner *notifications.NotificationService + available atomic.Bool +} + +func newNotifier() *Notifier { + return &Notifier{inner: notifications.New()} +} + +// ServiceName implements the Wails service-name hook for startup logs. +func (n *Notifier) ServiceName() string { + return n.inner.ServiceName() +} + +// ServiceStartup starts the platform notifier, downgrading failure to a +// warning so the app keeps running without notifications. +func (n *Notifier) ServiceStartup(ctx context.Context, options application.ServiceOptions) error { + if err := n.inner.ServiceStartup(ctx, options); err != nil { + log.Warnf("notifications disabled: %v", err) + return nil + } + n.available.Store(true) + return nil +} + +func (n *Notifier) ServiceShutdown() error { + if !n.available.Load() { + return nil + } + return n.inner.ServiceShutdown() +} + +func (n *Notifier) CheckNotificationAuthorization() (bool, error) { + if !n.available.Load() { + return false, errNotificationsUnavailable + } + return n.inner.CheckNotificationAuthorization() +} + +func (n *Notifier) RequestNotificationAuthorization() (bool, error) { + if !n.available.Load() { + return false, errNotificationsUnavailable + } + return n.inner.RequestNotificationAuthorization() +} + +// SendNotification delivers a notification, silently dropping it when the +// backend never started (notifications are best-effort everywhere). +func (n *Notifier) SendNotification(options notifications.NotificationOptions) error { + if !n.available.Load() { + log.Debugf("notifications disabled, dropping %q", options.ID) + return nil + } + return n.inner.SendNotification(options) +} + +func (n *Notifier) SendNotificationWithActions(options notifications.NotificationOptions) error { + if !n.available.Load() { + log.Debugf("notifications disabled, dropping %q", options.ID) + return nil + } + return n.inner.SendNotificationWithActions(options) +} + +func (n *Notifier) RegisterNotificationCategory(category notifications.NotificationCategory) error { + if !n.available.Load() { + return nil + } + return n.inner.RegisterNotificationCategory(category) +} + +// OnNotificationResponse registers the response callback. Pure Go state, so +// it is safe (and simply inert) when the backend never started. +// +//wails:ignore +func (n *Notifier) OnNotificationResponse(callback func(result notifications.NotificationResult)) { + n.inner.OnNotificationResponse(callback) +} diff --git a/client/ui/preferences/store.go b/client/ui/preferences/store.go index df6fbbb16..3b677016f 100644 --- a/client/ui/preferences/store.go +++ b/client/ui/preferences/store.go @@ -58,6 +58,10 @@ type UIPreferences struct { // decision has run for this OS user. It only ever transitions to true // and is never reset, so the default-on flow runs at most once, ever. AutostartInitialized bool `json:"autostartInitialized"` + // KeepConnectedOnQuit leaves the daemon connected when the GUI quits. + // Its false zero value preserves the historical disconnect-on-quit + // behaviour for preference files written before the field existed. + KeepConnectedOnQuit bool `json:"keepConnectedOnQuit"` } // LanguageValidator rejects SetLanguage inputs with no shipped bundle. @@ -183,6 +187,26 @@ func (s *Store) SetAutostartInitialized(done bool) error { return nil } +// SetKeepConnectedOnQuit persists the disconnect-on-quit opt-out. No-op if unchanged. +func (s *Store) SetKeepConnectedOnQuit(keep bool) error { + s.mu.Lock() + if s.current.KeepConnectedOnQuit == keep { + s.mu.Unlock() + return nil + } + next := s.current + next.KeepConnectedOnQuit = keep + if err := s.persistLocked(next); err != nil { + s.mu.Unlock() + return fmt.Errorf("persist preferences: %w", err) + } + s.current = next + s.mu.Unlock() + + s.broadcast(next) + return nil +} + // SetLanguage validates, persists, and broadcasts. No-op if unchanged. func (s *Store) SetLanguage(lang i18n.LanguageCode) error { if lang == "" { @@ -246,6 +270,7 @@ func (s *Store) ExistedAtLoad() bool { func (s *Store) load() error { if _, err := os.Stat(s.path); err != nil { if errors.Is(err, os.ErrNotExist) { + log.Infof("no ui preferences file at %s; using defaults", s.path) return nil } return fmt.Errorf("stat preferences: %w", err) diff --git a/client/ui/preferences/store_test.go b/client/ui/preferences/store_test.go index 6384fddb8..3e1cb3107 100644 --- a/client/ui/preferences/store_test.go +++ b/client/ui/preferences/store_test.go @@ -238,6 +238,42 @@ func TestStore_SetAutostartInitializedPersistsAcrossReload(t *testing.T) { assert.True(t, reloaded.Get().AutostartInitialized, "marker must survive a reload from disk") } +func TestStore_SetKeepConnectedOnQuitPersistsAcrossReload(t *testing.T) { + withTempConfigDir(t) + emitter := &recordingEmitter{} + s, err := NewStore(nil, emitter) + require.NoError(t, err) + + assert.False(t, s.Get().KeepConnectedOnQuit, "quitting must disconnect by default") + + require.NoError(t, s.SetKeepConnectedOnQuit(true)) + assert.True(t, s.Get().KeepConnectedOnQuit, "Get should reflect the persisted opt-out") + require.Len(t, emitter.calledWith(EventPreferencesChanged), 1, "first write should broadcast") + + require.NoError(t, s.SetKeepConnectedOnQuit(true)) + assert.Len(t, emitter.calledWith(EventPreferencesChanged), 1, "idempotent write should not broadcast again") + + reloaded, err := NewStore(nil, nil) + require.NoError(t, err) + assert.True(t, reloaded.Get().KeepConnectedOnQuit, "opt-out must survive a reload from disk") +} + +func TestStore_KeepConnectedOnQuitDefaultsFalseForPreExistingFile(t *testing.T) { + withTempConfigDir(t) + + // A preferences file written before the field existed must keep the + // historical disconnect-on-quit behaviour rather than silently opting out. + path, err := preferencesPath() + require.NoError(t, err) + require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755)) + require.NoError(t, os.WriteFile(path, []byte(`{"language":"en","viewMode":"default"}`), 0o600)) + + s, err := NewStore(nil, nil) + require.NoError(t, err) + assert.False(t, s.Get().KeepConnectedOnQuit, "a file predating the field must not opt out of disconnect-on-quit") + assert.True(t, s.ExistedAtLoad(), "the pre-existing file must be seen on disk") +} + func TestStore_ExistedAtLoad(t *testing.T) { withTempConfigDir(t) diff --git a/client/ui/services/autostart.go b/client/ui/services/autostart.go index f7e3aeea0..98e893f04 100644 --- a/client/ui/services/autostart.go +++ b/client/ui/services/autostart.go @@ -10,7 +10,7 @@ import ( "github.com/wailsapp/wails/v3/pkg/application" ) -// Autostart facade over Wails' AutostartManager. The OS login-item registration +// Autostart facade over Wails' AutostartManager. The OS autostart entry registration // is the single source of truth; nothing is mirrored to preferences. type Autostart struct { mgr *application.AutostartManager diff --git a/client/ui/services/connection.go b/client/ui/services/connection.go index fae7ddd23..f78ce4c0f 100644 --- a/client/ui/services/connection.go +++ b/client/ui/services/connection.go @@ -33,12 +33,21 @@ type LoginResult struct { UserCode string `json:"userCode"` VerificationURI string `json:"verificationUri"` VerificationURIComplete string `json:"verificationUriComplete"` + // ProfileID is the ID of the profile this login ran against, or "" when the + // caller named the profile itself and no ID was resolved. Pass it back in + // WaitSSOParams so the account email lands on this profile even if the + // active one changes during SSO. + ProfileID string `json:"profileId"` } // WaitSSOParams are the inputs to waitSSOLogin. type WaitSSOParams struct { UserCode string `json:"userCode"` Hostname string `json:"hostname"` + // ProfileID is the profile the login was started for, used to file the + // account email against it rather than against whichever profile is active + // when the flow returns. Optional: empty falls back to the active profile. + ProfileID string `json:"profileId"` } // UpParams selects the profile to bring up. @@ -77,11 +86,16 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err // Fall back to the daemon's active profile and the current OS user. profileName := p.ProfileName username := p.Username + // Only set when the daemon told us the ID. A caller-supplied ProfileName is + // a handle — a display name or an ID prefix resolve too — and the state file + // is named after the ID, so passing a handle on would name the wrong file. + profileID := "" if profileName == "" { if active, aerr := cli.GetActiveProfile(ctx, &proto.GetActiveProfileRequest{}); aerr == nil { // Address the active profile by ID (the daemon resolves it as a // handle); names can collide, the ID cannot. profileName = active.GetId() + profileID = profileName if username == "" { username = active.GetUsername() } @@ -94,10 +108,11 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err } req := &proto.LoginRequest{ - ManagementUrl: p.ManagementURL, - SetupKey: p.SetupKey, - Hostname: p.Hostname, - IsUnixDesktopClient: runtime.GOOS == "linux", + ManagementUrl: p.ManagementURL, + SetupKey: p.SetupKey, + Hostname: p.Hostname, + // a login driven by the UI always has a graphical session available + IsUnixDesktopClient: true, } if profileName != "" { req.ProfileName = ptrStr(profileName) @@ -108,8 +123,16 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err if p.PreSharedKey != "" { req.OptionalPreSharedKey = ptrStr(p.PreSharedKey) } - if p.Hint != "" { - req.Hint = ptrStr(p.Hint) + hint := p.Hint + if hint == "" && profileID != "" { + if state, serr := profilemanager.NewProfileManager().GetProfileState(profilemanager.ID(profileID)); serr == nil { + hint = state.Email + } else { + log.Debugf("failed to get profile state for login hint: %v", serr) + } + } + if hint != "" { + req.Hint = ptrStr(hint) } resp, err := cli.Login(ctx, req) @@ -122,6 +145,7 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err UserCode: resp.GetUserCode(), VerificationURI: resp.GetVerificationURI(), VerificationURIComplete: resp.GetVerificationURIComplete(), + ProfileID: profileID, }, nil } @@ -212,16 +236,6 @@ func (s *Connection) Logout(ctx context.Context, p LogoutParams) error { return s.classifyDaemonError(err) } - // The daemon runs as root and can't reach the user-owned per-profile state - // file holding the account email (see Profiles.List), so clear the stale - // email here; the next SSO login recreates it. - if p.ProfileName != "" { - if err := profilemanager.NewProfileManager().RemoveProfileState(p.ProfileName); err != nil { - // Non-fatal: the logout itself succeeded. - log.Warnf("failed to remove profile state for %s: %v", p.ProfileName, err) - } - } - return nil } @@ -242,6 +256,31 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string, return "", s.classifyDaemonError(err) } log.Infof("SSO login completed, daemon reported success") + + // Persist the account email the same way the CLI does after its own + // WaitSSOLogin: the daemon returns it but cannot store it, since it runs as + // root and the per-profile state file is user-owned (see Profiles.List). + // Without this the profile has no email, so Profiles.List shows no account + // and later logins and session extends go out without a login_hint — + // leaving the IdP to guess which account was meant. + if email := resp.GetEmail(); email != "" { + state := &profilemanager.ProfileState{Email: email} + pm := profilemanager.NewProfileManager() + + // Against the profile the login was started for: SSO spans seconds of + // user interaction, and a profile switch in that window would otherwise + // file the email under the wrong profile. + if p.ProfileID != "" { + err = pm.SetProfileState(profilemanager.ID(p.ProfileID), state) + } else { + err = pm.SetActiveProfileState(state) + } + if err != nil { + // Non-fatal: the login itself succeeded. + log.Warnf("failed to store account email: %v", err) + } + } + return resp.GetEmail(), nil } diff --git a/client/ui/services/cursor_linux.go b/client/ui/services/cursor_linux.go index 760fd5b86..3294f3c95 100644 --- a/client/ui/services/cursor_linux.go +++ b/client/ui/services/cursor_linux.go @@ -49,5 +49,11 @@ func getCursorPosition(app *application.App) (application.Point, bool) { if app == nil || app.Screen == nil { return p, true } + // The wails GTK3 backend caches screens from the active window; a tray app + // has none at startup, so the cache is empty and PhysicalToDipPoint would + // dereference a nil nearest screen. Raw pixels are correct there anyway. + if app.Screen.ScreenNearestPhysicalPoint(p) == nil { + return p, true + } return app.Screen.PhysicalToDipPoint(p), true } diff --git a/client/ui/services/debug.go b/client/ui/services/debug.go index 034086747..d1f6555a8 100644 --- a/client/ui/services/debug.go +++ b/client/ui/services/debug.go @@ -15,10 +15,13 @@ import ( ) type DebugBundleParams struct { - Anonymize bool `json:"anonymize"` - SystemInfo bool `json:"systemInfo"` - UploadURL string `json:"uploadUrl"` - LogFileCount uint32 `json:"logFileCount"` + Anonymize bool `json:"anonymize"` + // AnonymizeLevel is "default" or "strict"; strict also anonymizes + // private IP ranges, peer names, and WireGuard public keys. + AnonymizeLevel string `json:"anonymizeLevel"` + SystemInfo bool `json:"systemInfo"` + UploadURL string `json:"uploadUrl"` + LogFileCount uint32 `json:"logFileCount"` } // DebugBundleResult: Path is set for local-only bundles, UploadedKey on upload @@ -48,11 +51,12 @@ func (s *Debug) Bundle(ctx context.Context, p DebugBundleParams) (DebugBundleRes return DebugBundleResult{}, err } resp, err := cli.DebugBundle(ctx, &proto.DebugBundleRequest{ - Anonymize: p.Anonymize, - SystemInfo: p.SystemInfo, - UploadURL: p.UploadURL, - LogFileCount: p.LogFileCount, - CliVersion: version.NetbirdVersion(), + Anonymize: p.Anonymize, + AnonymizeLevel: p.AnonymizeLevel, + SystemInfo: p.SystemInfo, + UploadURL: p.UploadURL, + LogFileCount: p.LogFileCount, + CliVersion: version.NetbirdVersion(), }) if err != nil { return DebugBundleResult{}, err diff --git a/client/ui/services/errors.go b/client/ui/services/errors.go index f1679e764..0c6f2f20f 100644 --- a/client/ui/services/errors.go +++ b/client/ui/services/errors.go @@ -6,13 +6,30 @@ import ( "encoding/json" "strings" + "google.golang.org/genproto/googleapis/rpc/errdetails" gcodes "google.golang.org/grpc/codes" gstatus "google.golang.org/grpc/status" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/ui/i18n" "github.com/netbirdio/netbird/client/ui/preferences" ) +// privilegeErrorInfo returns the daemon's privilege-refusal detail, if the error +// carries one. +func privilegeErrorInfo(err error) (*errdetails.ErrorInfo, bool) { + for _, detail := range gstatus.Convert(err).Details() { + info, ok := detail.(*errdetails.ErrorInfo) + if !ok { + continue + } + if info.GetReason() == ipcauth.ErrorReasonPrivilegeRequired && info.GetDomain() == ipcauth.ErrorDomain { + return info, true + } + } + return nil, false +} + // ErrorTranslator localises daemon errors; runtime impl is *i18n.Bundle. type ErrorTranslator interface { Translate(lang i18n.LanguageCode, key string, args ...string) string @@ -30,6 +47,10 @@ type ClientError struct { Code string `json:"code"` Short string `json:"short"` Long string `json:"long"` + // Command is a command the user can run to complete the operation + // themselves, set when the daemon refused it for want of privileges. The + // frontend offers it for copying. + Command string `json:"command,omitempty"` } // Error returns the short message for plain Go callers. @@ -72,6 +93,24 @@ func (c errorClassifier) classify(err error) *ClientError { msg = st.Message() grpcCode = st.Code() } + + // A refusal for want of privileges carries its own summary and the command + // that performs the operation, both written for the user. Surface them + // verbatim: no substring guessing, and no localisation of a message the + // daemon composed. + if info, ok := privilegeErrorInfo(err); ok { + summary := info.GetMetadata()[ipcauth.ErrorMetaSummary] + if summary == "" { + summary = msg + } + return &ClientError{ + Code: "privilege_required", + Short: summary, + Long: summary, + Command: info.GetMetadata()[ipcauth.ErrorMetaCommand], + } + } + lower := strings.ToLower(msg) code := "unknown" diff --git a/client/ui/services/preferences.go b/client/ui/services/preferences.go index dae086de8..77faa4ef6 100644 --- a/client/ui/services/preferences.go +++ b/client/ui/services/preferences.go @@ -34,3 +34,7 @@ func (s *Preferences) SetViewMode(_ context.Context, mode preferences.ViewMode) func (s *Preferences) SetOnboardingCompleted(_ context.Context, done bool) error { return s.store.SetOnboardingCompleted(done) } + +func (s *Preferences) SetKeepConnectedOnQuit(_ context.Context, keep bool) error { + return s.store.SetKeepConnectedOnQuit(keep) +} diff --git a/client/ui/services/profile.go b/client/ui/services/profile.go index 09468c9df..e76ab3db6 100644 --- a/client/ui/services/profile.go +++ b/client/ui/services/profile.go @@ -6,6 +6,8 @@ import ( "context" "os/user" + log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/proto" ) @@ -151,11 +153,32 @@ func (s *Profiles) Remove(ctx context.Context, p ProfileRef) error { if err != nil { return err } - _, err = cli.RemoveProfile(ctx, &proto.RemoveProfileRequest{ + resp, err := cli.RemoveProfile(ctx, &proto.RemoveProfileRequest{ ProfileName: p.ProfileName, Username: p.Username, }) - return err + if err != nil { + return err + } + + // The daemon deletes what it owns but runs as root, so it leaves the + // user-owned state file holding the account email behind. Logout keeps the + // email on purpose so later logins can pass it as the login_hint; profile + // removal is what deletes it. Legacy profiles are keyed by name rather than by a + // generated ID, so a recreated profile of the same name would inherit the + // deleted one's email and offer it as the login_hint. + // + // Keyed on the ID the daemon resolved, not on the request handle: that may + // have been a display name or an ID prefix, which would name a different + // file (or none). + if id := resp.GetId(); id != "" { + if err := profilemanager.NewProfileManager().RemoveProfileState(id); err != nil { + // Non-fatal: the profile itself is gone. + log.Warnf("failed to remove profile state for %s: %v", id, err) + } + } + + return nil } // Rename changes a profile's display name. The on-disk ID is unaffected, so diff --git a/client/ui/services/profileswitcher.go b/client/ui/services/profileswitcher.go index c27b62d92..727b2473f 100644 --- a/client/ui/services/profileswitcher.go +++ b/client/ui/services/profileswitcher.go @@ -12,13 +12,15 @@ import ( "github.com/netbirdio/netbird/client/internal/profilemanager" ) -// ProfileSwitcher holds the reconnect policy shared by the tray and React -// frontend so both flip profiles identically. The policy keys off prevStatus -// from DaemonFeed.Get at SwitchActive entry: +// ProfileSwitcher holds the switch policy shared by the tray and React +// frontend so both flip profiles identically. SwitchActive (plain selection: +// header dropdown, tray submenu) always connects after the switch; +// SwitchActiveNoConnect (manage-profiles screen) never does, so the user can +// still adjust the management URL before connecting. prevStatus from +// DaemonFeed.Get at entry only decides the teardown: // -// Connected/Connecting → Switch + Down + Up; optimistic Connecting paint. -// NeedsLogin/LoginFailed/SessionExpired → Switch + Down; clear stale error for re-login. -// Idle → Switch only. +// Connected/Connecting/NeedsLogin/LoginFailed/SessionExpired → Down first. +// Idle → no Down. type ProfileSwitcher struct { profiles *Profiles connection *Connection @@ -29,29 +31,40 @@ func NewProfileSwitcher(profiles *Profiles, connection *Connection, feed *Daemon return &ProfileSwitcher{profiles: profiles, connection: connection, feed: feed} } -// SwitchActive switches to the named profile applying the reconnect policy. +// SwitchActive switches to the named profile and always connects afterwards. func (s *ProfileSwitcher) SwitchActive(ctx context.Context, p ProfileRef) error { + return s.switchActive(ctx, p, true) +} + +// SwitchActiveNoConnect switches to the named profile without connecting, +// tearing down any existing connection first. +func (s *ProfileSwitcher) SwitchActiveNoConnect(ctx context.Context, p ProfileRef) error { + return s.switchActive(ctx, p, false) +} + +func (s *ProfileSwitcher) switchActive(ctx context.Context, p ProfileRef, connect bool) error { prevStatus := "" - if st, err := s.feed.Get(ctx); err == nil { - prevStatus = st.Status - } else { - log.Warnf("profileswitcher: get status: %v", err) + if s.feed != nil { + if st, err := s.feed.Get(ctx); err == nil { + prevStatus = st.Status + } else { + log.Warnf("profileswitcher: get status: %v", err) + } } - wasActive := strings.EqualFold(prevStatus, StatusConnected) || - strings.EqualFold(prevStatus, StatusConnecting) - needsDown := wasActive || + needsDown := strings.EqualFold(prevStatus, StatusConnected) || + strings.EqualFold(prevStatus, StatusConnecting) || strings.EqualFold(prevStatus, StatusNeedsLogin) || strings.EqualFold(prevStatus, StatusLoginFailed) || strings.EqualFold(prevStatus, StatusSessionExpired) - log.Infof("profileswitcher: switch profile=%q prevStatus=%q wasActive=%v needsDown=%v", - p.ProfileName, prevStatus, wasActive, needsDown) + log.Infof("profileswitcher: switch profile=%q prevStatus=%q connect=%v needsDown=%v", + p.ProfileName, prevStatus, connect, needsDown) - // Optimistic Connecting paint only when wasActive: those prevStatuses emit - // stale Connected + transient Idle pushes during Down that must be - // suppressed until Up resumes the stream (see DaemonFeed suppression table). - if wasActive { + // Optimistic Connecting paint plus stale-push suppression during Down (see + // DaemonFeed suppression table); also arms the login-watch that pops + // browser-login when the new profile turns out to need SSO. + if connect && s.feed != nil { s.feed.BeginProfileSwitch() } @@ -76,9 +89,9 @@ func (s *ProfileSwitcher) SwitchActive(ctx context.Context, p ProfileRef) error } } - if wasActive { + if connect { if err := s.connection.Up(ctx, UpParams(p)); err != nil { - return fmt.Errorf("reconnect %q: %w", p.ProfileName, err) + return fmt.Errorf("connect %q: %w", p.ProfileName, err) } } diff --git a/client/ui/services/settings.go b/client/ui/services/settings.go index 3b6f6f81b..74e6f913c 100644 --- a/client/ui/services/settings.go +++ b/client/ui/services/settings.go @@ -7,6 +7,10 @@ import ( "fmt" "reflect" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" ) @@ -39,6 +43,19 @@ type Restrictions struct { Features Features `json:"features"` } +// Privilege tells the frontend whether this process may perform the changes the +// daemon restricts to root/administrator, and carries the command for each so a +// disabled control can show the way to do it. +type Privilege struct { + Privileged bool `json:"privileged"` + // Actor names what the operation requires ("root", "administrator privileges"). + Actor string `json:"actor"` + // Commands equivalent to the settings the daemon guards, ready to copy. + AllowSSHServer string `json:"allowSshServer"` + EnableSSHRoot string `json:"enableSshRoot"` + DisableSSHAuth string `json:"disableSshAuth"` +} + type ConfigParams struct { ProfileName string `json:"profileName"` Username string `json:"username"` @@ -106,11 +123,19 @@ type SetConfigParams struct { } type Settings struct { - conn DaemonConn + conn DaemonConn + classifier errorClassifier + // daemonAddr is where the daemon listens, used to tell whether it runs as + // this user and would therefore authorize us: see Privilege. + daemonAddr string } -func NewSettings(conn DaemonConn) *Settings { - return &Settings{conn: conn} +func NewSettings(conn DaemonConn, translator ErrorTranslator, prefs LanguagePreference, daemonAddr string) *Settings { + return &Settings{ + conn: conn, + classifier: errorClassifier{translator: translator, prefs: prefs}, + daemonAddr: daemonAddr, + } } func (s *Settings) GetConfig(ctx context.Context, p ConfigParams) (Config, error) { @@ -189,8 +214,47 @@ func (s *Settings) SetConfig(ctx context.Context, p SetConfigParams) error { DisableSSHAuth: p.DisableSSHAuth, SshJWTCacheTTL: p.SSHJWTCacheTTL, } - _, err = cli.SetConfig(ctx, req) - return err + if _, err := cli.SetConfig(ctx, req); err != nil { + // Classified so the frontend gets the daemon's guidance instead of the + // gRPC envelope, which is what a refused privileged change looks like. + return s.classifier.classify(err) + } + return nil +} + +// Privilege reports whether this UI process could carry out the changes the +// daemon restricts to root/administrator, and the command that performs the one +// users hit in the SSH settings. It applies the daemon's own rule to what it can +// see locally, so the frontend can present those controls as unavailable up front +// instead of letting a save fail. No daemon round-trip, so it also works while the +// daemon is down. +// +// Being root or an elevated administrator is one way. The other is running as the +// daemon's own user while the daemon is unprivileged, which the daemon accepts +// because such a caller can already rewrite the config it reads; that is the +// rootless-container and Windows netstack-mode case, and it is read from the +// ownership of the socket or pipe the daemon created. +func (s *Settings) Privilege() Privilege { + id, err := ipcauth.CurrentProcessIdentity() + if err != nil { + // Fail closed: report unprivileged, which only ever disables controls. + log.Warnf("cannot read this process's identity, treating it as unprivileged: %v", err) + return newPrivilege(false) + } + if id.IsPrivileged() { + return newPrivilege(true) + } + return newPrivilege(daemonaddr.DaemonRunsAsSelf(s.daemonAddr)) +} + +func newPrivilege(privileged bool) Privilege { + return Privilege{ + Privileged: privileged, + Actor: ipcauth.PrivilegedActor(), + AllowSSHServer: ipcauth.UpCommand("--allow-server-ssh"), + EnableSSHRoot: ipcauth.UpCommand("--enable-ssh-root"), + DisableSSHAuth: ipcauth.UpCommand("--disable-ssh-auth"), + } } func (s *Settings) GetRestrictions(ctx context.Context) (Restrictions, error) { diff --git a/client/ui/services/shutdown.go b/client/ui/services/shutdown.go new file mode 100644 index 000000000..0da51940c --- /dev/null +++ b/client/ui/services/shutdown.go @@ -0,0 +1,24 @@ +package services + +import "sync/atomic" + +var ( + sessionEnding atomic.Bool + quitting atomic.Bool +) + +func BeginSessionEnd() { + sessionEnding.Store(true) +} + +func AbortSessionEnd() { + sessionEnding.Store(false) +} + +func BeginShutdown() { + quitting.Store(true) +} + +func ShuttingDown() bool { + return sessionEnding.Load() || quitting.Load() +} diff --git a/client/ui/services/update.go b/client/ui/services/update.go index 753177d45..b743b9858 100644 --- a/client/ui/services/update.go +++ b/client/ui/services/update.go @@ -10,6 +10,7 @@ import ( "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/ui/updater" + "github.com/netbirdio/netbird/version" ) // UpdateResult mirrors TriggerUpdateResponse. @@ -33,6 +34,12 @@ func (s *Update) GetState() updater.State { return s.holder.Get() } +// DownloadURL returns the platform-appropriate installer download link for +// manual (non-enforced) updates. +func (s *Update) DownloadURL() string { + return version.DownloadUrl() +} + // Quit exits the app. Scheduled off the calling goroutine so the JS caller's // response returns before the runtime tears down. func (s *Update) Quit() { diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index 3316dadaa..4930ce22b 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -8,6 +8,7 @@ import ( "sync" "time" + log "github.com/sirupsen/logrus" "github.com/wailsapp/wails/v3/pkg/application" "github.com/wailsapp/wails/v3/pkg/events" @@ -29,6 +30,12 @@ const EventBrowserLoginCancel = "browser-login:cancel" // EventSettingsOpen tells the mounted settings window which tab to show. const EventSettingsOpen = "netbird:settings:open" +const EventWindowPainted = "netbird:window-painted" + +const paintedFallback = 2 * time.Second + +const headlessTeardownDelay = 2 * time.Second + var WindowBackgroundColour = application.NewRGB(24, 26, 29) // bg-nb-gray-950 // WindowHeight is shared by the main and Settings windows. @@ -94,9 +101,6 @@ func DialogWindowOptions(name, title, url string, linuxIcon []byte) application. } } -// WindowManager owns the auxiliary windows (main is created in main.go). Settings is created -// eagerly and hidden on close to keep React state; the rest are created on open, destroyed on -// close, so the macOS dock-reopen handler finds no hidden window to resurrect. type WindowManager struct { app *application.App mainWindow *application.WebviewWindow @@ -112,15 +116,35 @@ type WindowManager struct { // hiddenForLogin holds windows hidden while the BrowserLogin popup is open, restored on close. hiddenForLogin []application.Window mu sync.Mutex + createMu sync.Mutex + newMain func(startURL string) *application.WebviewWindow + ready map[uint]bool + showPending map[uint]bool + pendingTab map[uint]string + pendingEmits map[uint][]string + fallbackTimers map[uint]*time.Timer + headlessMain bool + headlessTimer *time.Timer // recenterOnShow is set only on the minimal-WM/XEmbed path, where the WM neither centers nor // restores position; nil on full desktops so re-centering can't fight a user-moved window. recenterOnShow func() bool } -// NewWindowManager wires the manager to the main app; translator/prefs may be nil (tests). The -// Settings window is created here (hidden) so the first OpenSettings is instant. func NewWindowManager(app *application.App, mainWindow *application.WebviewWindow, translator ErrorTranslator, prefs LanguagePreference, linuxIcon []byte) *WindowManager { - s := &WindowManager{app: app, mainWindow: mainWindow, translator: translator, prefs: prefs, linuxIcon: linuxIcon} + s := &WindowManager{ + app: app, + mainWindow: mainWindow, + translator: translator, + prefs: prefs, + linuxIcon: linuxIcon, + ready: map[uint]bool{}, + showPending: map[uint]bool{}, + pendingTab: map[uint]string{}, + pendingEmits: map[uint][]string{}, + fallbackTimers: map[uint]*time.Timer{}, + } + s.watchPainted() + s.watchTriggerLogin() // Re-title live windows on language flip. Wired internally so the binding generator // doesn't try to expose the interface param. if sub, ok := prefs.(LanguageSubscriber); ok && sub != nil { @@ -136,7 +160,11 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo } }() } - s.settings = app.Window.NewWithOptions(application.WebviewWindowOptions{ + return s +} + +func (s *WindowManager) newSettingsWindow() *application.WebviewWindow { + w := s.app.Window.NewWithOptions(application.WebviewWindowOptions{ Name: "settings", Title: s.title("window.title.settings"), Width: 900, @@ -150,15 +178,15 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo URL: "/#/settings", Mac: AppleMacOSAppearanceOptions(), Windows: MicrosoftWindowsAppearanceOptions(), - Linux: LinuxAppearanceOptions(linuxIcon), + Linux: LinuxAppearanceOptions(s.linuxIcon), }) - // Hide (not destroy) on close to keep React state; reset to General for a flash-free reopen. - s.settings.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) { - e.Cancel() - s.app.Event.Emit(EventSettingsOpen, "general") - s.settings.Hide() + w.RegisterHook(events.Common.WindowClosing, func(_ *application.WindowEvent) { + s.mu.Lock() + s.settings = nil + s.forgetWindowLocked(w) + s.mu.Unlock() }) - return s + return w } // OpenSettings shows the settings window on tab (empty → General), switching tab via @@ -168,11 +196,20 @@ func (s *WindowManager) OpenSettings(tab string) { if target == "" { target = "general" } - s.app.Event.Emit(EventSettingsOpen, target) - s.settings.Show() - s.settings.Focus() - // Re-center (minimal-WM only; see centerWhenReady). - s.centerWhenReady(s.settings) + + w, _ := s.ensureWindow(&s.settings, s.newSettingsWindow) + + s.mu.Lock() + ready := s.ready[w.ID()] + if !ready { + s.pendingTab[w.ID()] = target + } + s.mu.Unlock() + + if ready { + s.app.Event.Emit(EventSettingsOpen, target) + } + s.showWhenReady(w) } // OpenBrowserLogin shows the SSO popup, creating it on first use. @@ -185,37 +222,38 @@ func (s *WindowManager) OpenBrowserLogin(uri string) { startURL = "/#/dialog/browser-login?uri=" + url.QueryEscape(uri) } s.hideOtherWindowsLocked("browser-login") - // Prefer the main window's screen (multi-monitor); falls back to OS-default centering. - var screen *application.Screen - if s.mainWindow != nil { - if sc, err := s.mainWindow.GetScreen(); err == nil { - screen = sc - } - } opts := DialogWindowOptions("browser-login", s.title("window.title.signIn"), startURL, s.linuxIcon) // Not always-on-top: it would obscure the browser tab the user logs in through. opts.AlwaysOnTop = false opts.InitialPosition = application.WindowCentered - opts.Screen = screen + // Open on the active (where users cursor is) display, like the session-expiration dialog. + opts.Screen = s.getScreenBasedOnCursorPosition() s.browserLogin = s.app.Window.NewWithOptions(opts) bl := s.browserLogin - // Red-X close means cancel: emit the event so startLogin() tears down the SSO wait. bl.OnWindowEvent(events.Common.WindowClosing, func(_ *application.WindowEvent) { - s.app.Event.Emit(EventBrowserLoginCancel) s.mu.Lock() - s.browserLogin = nil - s.restoreHiddenWindowsLocked() + // Only a live user red-X still has this registered; programmatic closers + // nil s.browserLogin first and clean up themselves. Guarding here stops a + // stale close event from wiping a replacement popup's state. + userClosed := s.browserLogin == bl + if userClosed { + s.browserLogin = nil + s.restoreHiddenWindowsLocked() + } s.mu.Unlock() + if userClosed { + s.app.Event.Emit(EventBrowserLoginCancel) + } }) - s.centerWhenReady(s.browserLogin) + s.centerOnCursorScreen(s.browserLogin) return } if uri != "" { s.browserLogin.SetURL("/#/dialog/browser-login?uri=" + url.QueryEscape(uri)) } + s.centerOnCursorScreen(s.browserLogin) s.browserLogin.Show() s.browserLogin.Focus() - s.centerWhenReady(s.browserLogin) } // BrowserLoginWindow returns the live SSO popup, or nil. While non-nil it is the @@ -238,6 +276,15 @@ func (s *WindowManager) CloseBrowserLogin() { s.mu.Lock() w := s.browserLogin s.browserLogin = nil + // 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. + if w != nil { + s.restoreHiddenWindowsLocked() + } s.mu.Unlock() if w != nil { w.Close() @@ -279,6 +326,35 @@ func (s *WindowManager) CloseSessionExpiration() { } } +// CloseRenewFlow tears down the SSO session-renewal UI in a single call: it +// closes the browser-login popup and the session-expiration window together. +func (s *WindowManager) CloseRenewFlow() { + s.mu.Lock() + bl := s.browserLogin + se := s.sessionExpiration + s.browserLogin = nil + s.sessionExpiration = nil + if se != nil { + kept := s.hiddenForLogin[:0] + for _, w := range s.hiddenForLogin { + if w != se { + kept = append(kept, w) + } + } + s.hiddenForLogin = kept + } + s.restoreHiddenWindowsLocked() + s.mu.Unlock() + + // Close after unlock so the re-entrant handlers can take s.mu. + if bl != nil { + bl.Close() + } + if se != nil { + se.Close() + } +} + // OpenInstallProgress shows the install-progress window and hides the rest for the duration // (restored on close). It owns its own result polling since the daemon restarts mid-install. func (s *WindowManager) OpenInstallProgress(version string) { @@ -351,12 +427,17 @@ func (s *WindowManager) CloseWelcome() { } } -// OpenError shows the custom error dialog; title/message are pre-localised and ride in the -// start URL. A second error replaces the open one via SetURL. Singleton, destroyed on close. -func (s *WindowManager) OpenError(title, message string) { +// OpenError shows the custom error dialog; title/message/command are pre-localised +// and ride in the start URL. command is optional and, when set, is offered for +// copying so the user can run the operation the daemon refused. A second error +// replaces the open one via SetURL. Singleton, destroyed on close. +func (s *WindowManager) OpenError(title, message, command string) { + if ShuttingDown() { + return + } s.mu.Lock() defer s.mu.Unlock() - startURL := errorDialogURL(title, message) + startURL := errorDialogURL(title, message, command) if s.errorDialog == nil { s.errorDialog = s.app.Window.NewWithOptions( DialogWindowOptions("error", s.title("window.title.error"), startURL, s.linuxIcon), @@ -393,13 +474,295 @@ func (s *WindowManager) OpenMain() { // ShowMain brings the main window forward (re-centering on minimal WMs). The single entry // point every surface (tray, SIGUSR1, welcome) should use so centering applies uniformly. func (s *WindowManager) ShowMain() { - if s.mainWindow == nil { + s.showWhenReady(s.MainWindow()) +} + +// ShowMainAndEmit brings the main window forward and emits event once its frontend is ready. +func (s *WindowManager) ShowMainAndEmit(event string) { + w := s.MainWindow() + if w == nil { return } - s.mainWindow.Show() - s.mainWindow.Focus() - // Re-center (minimal-WM only; see centerWhenReady). - s.centerWhenReady(s.mainWindow) + + id := w.ID() + s.mu.Lock() + ready := s.ready[id] + if !ready { + s.pendingEmits[id] = append(s.pendingEmits[id], event) + } + s.mu.Unlock() + + s.showWhenReady(w) + if ready { + s.app.Event.Emit(event) + } +} + +func (s *WindowManager) MainWindow() *application.WebviewWindow { + w, _ := s.ensureMain("/") + return w +} + +func (s *WindowManager) ensureMain(startURL string) (*application.WebviewWindow, bool) { + s.mu.Lock() + factory := s.newMain + s.mu.Unlock() + if factory == nil { + return s.ensureWindow(&s.mainWindow, nil) + } + return s.ensureWindow(&s.mainWindow, func() *application.WebviewWindow { + return factory(startURL) + }) +} + +func (s *WindowManager) ensureWindow(slot **application.WebviewWindow, factory func() *application.WebviewWindow) (*application.WebviewWindow, bool) { + s.createMu.Lock() + defer s.createMu.Unlock() + + s.mu.Lock() + w := *slot + s.mu.Unlock() + if w != nil || factory == nil { + return w, false + } + + w = factory() + s.armReady(w) + + s.mu.Lock() + *slot = w + s.mu.Unlock() + return w, true +} + +func (s *WindowManager) armReady(w *application.WebviewWindow) { + if w == nil { + return + } + 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() + }) +} + +func (s *WindowManager) watchPainted() { + s.app.Event.On(EventWindowPainted, func(e *application.CustomEvent) { + if w := s.windowByName(e.Sender); w != nil { + s.markReady(w) + } + }) +} + +func (s *WindowManager) watchTriggerLogin() { + s.app.Event.On(EventTriggerLogin, func(_ *application.CustomEvent) { + s.mu.Lock() + if s.headlessTimer != nil { + s.headlessTimer.Stop() + s.headlessTimer = nil + } + w := s.mainWindow + ready := w != nil && s.ready[w.ID()] + s.mu.Unlock() + if ready { + return + } + + w, created := s.ensureMain("/") + if w == nil { + return + } + + s.mu.Lock() + if created { + s.headlessMain = true + } + pending := !s.ready[w.ID()] + if pending { + s.pendingEmits[w.ID()] = append(s.pendingEmits[w.ID()], EventTriggerLogin) + } + s.mu.Unlock() + + if !pending { + s.app.Event.Emit(EventTriggerLogin) + } + }) + + s.app.Event.On(EventBrowserLoginCancel, func(_ *application.CustomEvent) { + s.scheduleHeadlessTeardown() + }) + + s.app.Event.On(EventStatusSnapshot, func(e *application.CustomEvent) { + st, ok := e.Data.(Status) + if !ok { + return + } + switch st.Status { + case StatusConnected, StatusLoginFailed, StatusDaemonUnavailable: + s.scheduleHeadlessTeardown() + } + }) +} + +func (s *WindowManager) scheduleHeadlessTeardown() { + s.mu.Lock() + defer s.mu.Unlock() + if !s.headlessMain || s.mainWindow == nil { + return + } + if s.headlessTimer != nil { + s.headlessTimer.Stop() + } + s.headlessTimer = time.AfterFunc(headlessTeardownDelay, s.closeHeadlessMain) +} + +func (s *WindowManager) closeHeadlessMain() { + s.mu.Lock() + w := s.mainWindow + headless := s.headlessMain + s.headlessTimer = nil + s.mu.Unlock() + if !headless || w == nil { + return + } + w.Close() +} + +func (s *WindowManager) forgetWindowLocked(w *application.WebviewWindow) { + if w == nil { + return + } + + id := w.ID() + if timer := s.fallbackTimers[id]; timer != nil { + timer.Stop() + } + delete(s.fallbackTimers, id) + delete(s.ready, id) + delete(s.showPending, id) + delete(s.pendingTab, id) + delete(s.pendingEmits, id) + + kept := s.hiddenForLogin[:0] + for _, hidden := range s.hiddenForLogin { + if hidden != application.Window(w) { + kept = append(kept, hidden) + } + } + s.hiddenForLogin = kept +} + +func (s *WindowManager) windowByName(name string) *application.WebviewWindow { + s.mu.Lock() + defer s.mu.Unlock() + switch name { + case "main": + return s.mainWindow + case "settings": + return s.settings + default: + return nil + } +} + +func (s *WindowManager) markReady(w *application.WebviewWindow) { + id := w.ID() + s.mu.Lock() + already := s.ready[id] + s.ready[id] = true + wanted := s.showPending[id] + tab, hasTab := s.pendingTab[id] + emits := s.pendingEmits[id] + if timer := s.fallbackTimers[id]; timer != nil { + timer.Stop() + delete(s.fallbackTimers, id) + } + delete(s.showPending, id) + delete(s.pendingTab, id) + delete(s.pendingEmits, id) + s.mu.Unlock() + + if already { + return + } + + if hasTab { + s.app.Event.Emit(EventSettingsOpen, tab) + } + + if wanted { + s.showNow(w) + } + + for _, event := range emits { + s.app.Event.Emit(event) + } +} + +func (s *WindowManager) showWhenReady(w *application.WebviewWindow) { + if w == nil { + return + } + + id := w.ID() + s.mu.Lock() + ready := s.ready[id] + if !ready { + s.showPending[id] = true + } + s.mu.Unlock() + + if ready { + s.showNow(w) + } +} + +func (s *WindowManager) showNow(w *application.WebviewWindow) { + s.mu.Lock() + if w == s.mainWindow { + s.headlessMain = false + if s.headlessTimer != nil { + s.headlessTimer.Stop() + s.headlessTimer = nil + } + } + s.mu.Unlock() + w.Show() + w.Focus() + s.centerWhenReady(w) +} + +func (s *WindowManager) ShowMainAt(url string) { + w, created := s.ensureMain(url) + if w == nil { + return + } + if !created { + w.SetURL(url) + } + s.showWhenReady(w) +} + +func (s *WindowManager) SetMainFactory(f func(startURL string) *application.WebviewWindow) { + s.mu.Lock() + defer s.mu.Unlock() + s.newMain = f +} + +func (s *WindowManager) ForgetMain() { + s.mu.Lock() + defer s.mu.Unlock() + s.forgetWindowLocked(s.mainWindow) + s.mainWindow = nil + s.headlessMain = false + if s.headlessTimer != nil { + s.headlessTimer.Stop() + s.headlessTimer = nil + } } // SetRecenterOnShow installs the recenterOnShow predicate (see the field). @@ -556,8 +919,8 @@ func (s *WindowManager) getScreenBasedOnCursorPosition() *application.Screen { return nil } -// errorDialogURL builds the error window's start URL with title/message as escaped query params. -func errorDialogURL(title, message string) string { +// errorDialogURL builds the error window's start URL with title/message/command as escaped query params. +func errorDialogURL(title, message, command string) string { q := url.Values{} if title != "" { q.Set("title", title) @@ -565,6 +928,9 @@ func errorDialogURL(title, message string) string { if message != "" { q.Set("message", message) } + if command != "" { + q.Set("command", command) + } startURL := "/#/dialog/error" if enc := q.Encode(); enc != "" { startURL += "?" + enc diff --git a/client/ui/services/windowtheme_windows.go b/client/ui/services/windowtheme_windows.go new file mode 100644 index 000000000..7dbc1164b --- /dev/null +++ b/client/ui/services/windowtheme_windows.go @@ -0,0 +1,14 @@ +package services + +import "github.com/wailsapp/wails/v3/pkg/w32" + +// Wails assigns w32.AllowDarkModeForWindow only on builds >= 18334 but calls it +// without a nil check when a window requests the Dark theme, crashing older +// builds such as Windows Server 2019 (17763). Those builds still get a dark +// title bar via the pre-20H1 DWM attribute that w32.SetTheme applies, so a +// no-op stub keeps the Dark theme fully working there. +func init() { + if w32.AllowDarkModeForWindow == nil { + w32.AllowDarkModeForWindow = func(w32.HWND, bool) uintptr { return 0 } + } +} diff --git a/client/ui/shutdown_other.go b/client/ui/shutdown_other.go new file mode 100644 index 000000000..6e617233f --- /dev/null +++ b/client/ui/shutdown_other.go @@ -0,0 +1,7 @@ +//go:build !windows && !android && !ios && !freebsd && !js + +package main + +func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) { + return nil +} diff --git a/client/ui/shutdown_windows.go b/client/ui/shutdown_windows.go new file mode 100644 index 000000000..fbb92a518 --- /dev/null +++ b/client/ui/shutdown_windows.go @@ -0,0 +1,36 @@ +//go:build windows + +package main + +import ( + "os" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/ui/services" +) + +const ( + wmQueryEndSession = 0x0011 + wmEndSession = 0x0016 +) + +func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) { + return func(_ uintptr, msg uint32, wParam, _ uintptr) (uintptr, bool) { + switch msg { + case wmQueryEndSession: + services.BeginSessionEnd() + return 1, true + case wmEndSession: + if wParam == 0 { + services.AbortSessionEnd() + return 0, true + } + log.Info("windows session is ending; exiting immediately") + os.Exit(0) + return 0, true + default: + return 0, false + } + } +} diff --git a/client/ui/tray.go b/client/ui/tray.go index 700d94098..c392a0b62 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -16,6 +16,7 @@ import ( "github.com/netbirdio/netbird/client/ui/authsession" "github.com/netbirdio/netbird/client/ui/i18n" + "github.com/netbirdio/netbird/client/ui/preferences" "github.com/netbirdio/netbird/client/ui/services" "github.com/netbirdio/netbird/version" ) @@ -30,9 +31,10 @@ const ( statusError = "Error" - urlGitHubRepo = "https://github.com/netbirdio/netbird" - urlGitHubReleases = "https://github.com/netbirdio/netbird/releases/latest" - urlDocs = "https://docs.netbird.io" + quitDownTimeout = 5 * time.Second + + urlGitHubRepo = "https://github.com/netbirdio/netbird" + urlDocs = "https://docs.netbird.io" ) // TrayServices bundles the services the tray menu needs, grouped so NewTray @@ -43,14 +45,15 @@ type TrayServices struct { Profiles *services.Profiles Networks *services.Networks DaemonFeed *services.DaemonFeed - Notifier *notifications.NotificationService + Notifier *Notifier Update *services.Update ProfileSwitcher *services.ProfileSwitcher WindowManager *services.WindowManager // Session is bound to authsession directly because the services wrapper // only re-exposes the React subset. - Session *authsession.Session - Localizer *Localizer + Session *authsession.Session + Localizer *Localizer + Preferences *preferences.Store } type Tray struct { @@ -171,7 +174,7 @@ func NewTray(app *application.App, window *application.WebviewWindow, svc TraySe // in the right locale — no English flash then re-paint. loc: svc.Localizer, } - t.updater = newTrayUpdater(app, window, svc.Update, svc.Notifier, t.loc, func() { t.applyIcon() }, func() { t.relayoutMenu() }) + t.updater = newTrayUpdater(app, t.showMainAt, svc.Update, svc.Notifier, t.loc, func() { t.applyIcon() }, func() { t.relayoutMenu() }) t.tray = app.SystemTray.New() // Seed panel-theme detection before the first paint so the initial icon // matches the panel's light/dark scheme (Linux only). @@ -238,9 +241,6 @@ func (t *Tray) ShowWindow() { w.Focus() return } - if t.window == nil { - return - } // Route through WindowManager so the main window is centered on first // show — minimal WMs (fluxbox, the XEmbed tray path) otherwise drop it in // the top-left corner. @@ -248,8 +248,49 @@ func (t *Tray) ShowWindow() { t.svc.WindowManager.ShowMain() return } - t.window.Show() - t.window.Focus() + if w := t.mainWindow(); w != nil { + w.Show() + w.Focus() + } +} + +func (t *Tray) mainWindow() *application.WebviewWindow { + if t.svc.WindowManager == nil { + return t.window + } + return t.svc.WindowManager.MainWindow() +} + +func (t *Tray) showMainAt(url string) { + if t.svc.WindowManager != nil { + t.svc.WindowManager.ShowMainAt(url) + return + } + if w := t.mainWindow(); w != nil { + w.SetURL(url) + w.Show() + w.Focus() + } +} + +func (t *Tray) showMain() { + if t.svc.WindowManager != nil { + t.svc.WindowManager.ShowMain() + return + } + if w := t.mainWindow(); w != nil { + w.Show() + w.Focus() + } +} + +func (t *Tray) showMainAndEmit(event string) { + if t.svc.WindowManager != nil { + t.svc.WindowManager.ShowMainAndEmit(event) + return + } + t.showMain() + t.app.Event.Emit(event) } // applyLanguage re-renders every translated surface in the Localizer's current @@ -315,8 +356,7 @@ func (t *Tray) relayoutMenu() { if sessionDeadline.IsZero() { t.sessionExpiresItem.SetHidden(true) } else { - remaining := t.formatSessionRemaining(time.Until(sessionDeadline)) - t.sessionExpiresItem.SetLabel(t.loc.T("tray.session.expiresIn", "remaining", remaining)) + t.sessionExpiresItem.SetLabel(t.sessionRowLabel(sessionDeadline)) t.sessionExpiresItem.SetHidden(false) } } @@ -446,18 +486,39 @@ func (t *Tray) buildMenu() *application.Menu { menu.AddSeparator() menu.Add(t.loc.T("tray.menu.quit")). SetAccelerator("CmdOrCtrl+Q"). - OnClick(func(*application.Context) { t.app.Quit() }) + OnClick(func(*application.Context) { t.handleQuit() }) return menu } +func (t *Tray) handleQuit() { + services.BeginShutdown() + t.profileMu.Lock() + if t.switchCancel != nil { + t.switchCancel() + t.switchCancel = nil + } + t.profileMu.Unlock() + t.svc.DaemonFeed.CancelProfileSwitch() + + if t.svc.Preferences == nil || !t.svc.Preferences.Get().KeepConnectedOnQuit { + ctx, cancel := context.WithTimeout(context.Background(), quitDownTimeout) + defer cancel() + if err := t.svc.Connection.Down(ctx); err != nil { + log.Errorf("disconnect on quit: %v", err) + } + } + t.app.Quit() +} + // handleConnect receives the clicked item from the buildMenu closure — // t.upItem is menuMu-guarded and must not be read here. func (t *Tray) handleConnect(upItem *application.MenuItem) { // NeedsLogin/SessionExpired/LoginFailed won't honor a plain Up RPC — they // need the Login → WaitSSOLogin → Up sequence. Emit EventTriggerLogin so // the React startLogin() (which owns the BrowserLogin popup) drives it; - // the hidden main webview is alive and subscribed, so only the popup shows. + // the WindowManager materialises a hidden main webview when none is live, + // so only the popup shows. t.statusMu.Lock() needsLogin := strings.EqualFold(t.lastStatus, services.StatusNeedsLogin) || strings.EqualFold(t.lastStatus, services.StatusSessionExpired) || diff --git a/client/ui/tray_click_linux.go b/client/ui/tray_click_linux.go index 34a364fd9..95f5dfe85 100644 --- a/client/ui/tray_click_linux.go +++ b/client/ui/tray_click_linux.go @@ -4,17 +4,26 @@ package main // bindTrayClick wires the tray icon's left-click handler on Linux. // -// Both Linux click paths converge on Wails' linuxSystemTray.Activate, which -// fires the registered clickHandler: -// - Real SNI hosts (KDE Plasma, Waybar, GNOME Shell + AppIndicator) invoke -// org.kde.StatusNotifierItem.Activate over D-Bus on left-click. -// - The in-process StatusNotifierWatcher + XEmbed host used on minimal WMs -// (Fluxbox, i3, dwm, OpenBox) maps a Button1 press to that same Activate -// call itself (xembed_host_linux.go), so it routes through the same hook. -// Registering OnClick here therefore covers both paths with one handler — no -// changes to the watcher or XEmbed C code are needed. Left-click now opens the -// main window; right-click still opens the menu via Wails' default -// SecondaryActivate→OpenMenu handler (and the XEmbed GTK popup on minimal WMs). +// Expected behaviour per tray host: +// +// Host Left click Right click +// KDE Plasma, Waybar main window (Activate) menu (host-rendered) +// GNOME Shell + AppIndicator menu only menu only +// Minimal WMs via XEmbed host main window (Activate) XEmbed GTK popup +// +// OnClick fires only on org.kde.StatusNotifierItem.Activate — a real left +// click. KDE/Waybar send it over D-Bus; the in-process XEmbed host +// (xembed_host_linux.go) maps a Button1 press to the same Activate call. +// +// GNOME Shell + AppIndicator never sends Activate: it renders the dbusmenu +// on ANY click and only reports the menu opening via dbusmenu +// Event("opened"). Upstream Wails treated that event as a click, so on GNOME +// both buttons raised the main window on top of the menu, and on KDE/Waybar +// a right click raised it over the freshly opened menu. The netbirdio/wails +// fork (go.mod replace) drops that heuristic: a menu open never fires +// OnClick. On GNOME the main window is reached via the "Open NetBird" menu +// entry; left-click-opens-window is not achievable there anyway, since the +// host always opens the menu itself. // // We do NOT register OnDoubleClick: Wails' Linux SNI backend never fires it // (unlike Windows). And we deliberately skip AttachWindow — it plus Wails3's diff --git a/client/ui/tray_notify.go b/client/ui/tray_notify.go index 6c60e3d4b..5b2629419 100644 --- a/client/ui/tray_notify.go +++ b/client/ui/tray_notify.go @@ -25,6 +25,9 @@ type sendFn func(notifications.NotificationOptions) error // event-dispatch goroutine that panic is fatal process-wide; recover() turns // it into a logged no-op. func safeSendNotification(send sendFn, what string, opts notifications.NotificationOptions) (err error) { + if services.ShuttingDown() { + return nil + } defer func() { if r := recover(); r != nil { log.Errorf("notify %s: recovered from panic (notification bus unavailable): %v", what, r) @@ -41,7 +44,7 @@ func safeSendNotification(send sendFn, what string, opts notifications.Notificat // notifyIfDaemonOutdated probes the daemon once and fires an OS toast when it // is reachable but too old for this UI. A probe error means the daemon isn't // reachable (not outdated), so it is left to the normal connection flow. -func notifyIfDaemonOutdated(compat *services.Compat, notifier *notifications.NotificationService, loc *Localizer) { +func notifyIfDaemonOutdated(compat *services.Compat, notifier *Notifier, loc *Localizer) { ready, err := compat.DaemonReady(context.Background()) if err != nil { log.Debugf("daemon compatibility probe: %v", err) diff --git a/client/ui/tray_session.go b/client/ui/tray_session.go index 6b73ddb49..6e5d07740 100644 --- a/client/ui/tray_session.go +++ b/client/ui/tray_session.go @@ -27,14 +27,10 @@ const ( finalWarningCountdownSeconds = 120 ) -// handleSessionExpired notifies and brings the window forward so the frontend's /login route drives renewal. +// handleSessionExpired notifies and brings the window forward so the user can reconnect. func (t *Tray) handleSessionExpired() { t.notify(t.loc.T("notify.sessionExpired.title"), t.loc.T("notify.sessionExpired.body"), notifyIDSessionExpired) - if t.window != nil { - t.window.SetURL("/#/login") - t.window.Show() - t.window.Focus() - } + t.showMain() } // applySessionExpiry refreshes the cached SSO deadline and reports whether it changed. @@ -64,11 +60,42 @@ func (t *Tray) applySessionExpiry(deadline *time.Time, connected bool) bool { return changed } -// runSessionExpiryTicker recomputes the "Expires in …" row label every 30s. Runs until process exit. +// runSessionExpiryTicker recomputes the "Expires in …" row label until process exit. +// The interval scales with the remaining time: coarse when the deadline is far off, +// down to 10s in the final two minutes so the label doesn't lag the ceiling-rounded +// countdown near expiry. The cached deadline is re-read every iteration, so an extend +// or reconnect that moves it is picked up on the next tick. func (t *Tray) runSessionExpiryTicker() { - tk := time.NewTicker(30 * time.Second) - for range tk.C { + tm := time.NewTimer(sessionRefreshInterval(t.sessionRemaining())) + defer tm.Stop() + for range tm.C { t.refreshSessionExpiresLabel() + tm.Reset(sessionRefreshInterval(t.sessionRemaining())) + } +} + +// sessionRemaining returns the time left on the cached SSO deadline, or 0 when unknown. +func (t *Tray) sessionRemaining() time.Duration { + t.sessionMu.Lock() + deadline := t.sessionExpiresAt + t.sessionMu.Unlock() + if deadline.IsZero() { + return 0 + } + return time.Until(deadline) +} + +// sessionRefreshInterval picks how long to wait before the next label recompute. +func sessionRefreshInterval(remaining time.Duration) time.Duration { + switch { + case remaining <= 0: + return 30 * time.Second + case remaining <= 2*time.Minute: + return 10 * time.Second + case remaining <= time.Hour: + return 30 * time.Second + default: + return time.Minute } } @@ -87,30 +114,39 @@ func (t *Tray) refreshSessionExpiresLabel() { if deadline.IsZero() { return } - remaining := t.formatSessionRemaining(time.Until(deadline)) - item.SetLabel(t.loc.T("tray.session.expiresIn", "remaining", remaining)) + item.SetLabel(t.sessionRowLabel(deadline)) +} + +func (t *Tray) sessionRowLabel(deadline time.Time) string { + remaining := time.Until(deadline) + if remaining <= 0 { + return t.loc.T("tray.status.sessionExpired") + } + return t.loc.T("tray.session.expiresIn", "remaining", t.formatSessionRemaining(remaining)) } // formatSessionRemaining renders d as a localised long-form string picking the largest non-zero unit. +// Each unit is rounded up so the label never claims less time than actually remains, matching the +// upper-bound sense of the sub-minute "less than a minute" fragment. // Singular/plural keys are split per language for proper translation. func (t *Tray) formatSessionRemaining(d time.Duration) string { switch { case d < time.Minute: return t.loc.T("tray.session.unit.lessThanMinute") - case d < time.Hour: - m := int(d / time.Minute) + case d <= 59*time.Minute: + m := ceilDiv(d, time.Minute) if m == 1 { return t.loc.T("tray.session.unit.minute") } return t.loc.T("tray.session.unit.minutes", "count", strconv.Itoa(m)) - case d < 24*time.Hour: - h := int((d + 30*time.Minute) / time.Hour) + case d <= 23*time.Hour: + h := ceilDiv(d, time.Hour) if h == 1 { return t.loc.T("tray.session.unit.hour") } return t.loc.T("tray.session.unit.hours", "count", strconv.Itoa(h)) default: - days := int((d + 12*time.Hour) / (24 * time.Hour)) + days := ceilDiv(d, 24*time.Hour) if days == 1 { return t.loc.T("tray.session.unit.day") } @@ -118,6 +154,11 @@ func (t *Tray) formatSessionRemaining(d time.Duration) string { } } +// ceilDiv divides d by unit rounding up, assuming d > 0. +func ceilDiv(d, unit time.Duration) int { + return int((d + unit - time.Nanosecond) / unit) +} + // registerSessionWarningCategory wires the OS notification category and response handler for the expiry warning. // Errors are swallowed since the worst case is a plain notification without buttons. func (t *Tray) registerSessionWarningCategory() { @@ -252,11 +293,9 @@ func (t *Tray) openSessionExpiration() { } // openSessionExtendFlow opens the SessionExpiration window seeded with the cached deadline's remaining time, -// for the "Expires in …" tray row. No-ops when the deadline is unknown or elapsed. +// for the "Expires in …" tray row. Once the deadline has elapsed the row reads "Session expired" and the +// click routes to the login flow instead. No-op when the deadline is unknown. func (t *Tray) openSessionExtendFlow() { - if t.svc.WindowManager == nil { - return - } t.sessionMu.Lock() deadline := t.sessionExpiresAt t.sessionMu.Unlock() @@ -265,6 +304,10 @@ func (t *Tray) openSessionExtendFlow() { } seconds := int(time.Until(deadline).Seconds()) if seconds <= 0 { + t.showMainAndEmit(services.EventTriggerLogin) + return + } + if t.svc.WindowManager == nil { return } t.svc.WindowManager.OpenSessionExpiration(seconds) diff --git a/client/ui/tray_update.go b/client/ui/tray_update.go index 2d79cff05..3ce1f9600 100644 --- a/client/ui/tray_update.go +++ b/client/ui/tray_update.go @@ -4,6 +4,7 @@ package main import ( "context" + neturl "net/url" "sync" "time" @@ -13,14 +14,15 @@ import ( "github.com/netbirdio/netbird/client/ui/services" "github.com/netbirdio/netbird/client/ui/updater" + "github.com/netbirdio/netbird/version" ) // trayUpdater owns the tray UI that reacts to auto-update. Composed inside Tray. type trayUpdater struct { app *application.App - window *application.WebviewWindow + showMainAt func(url string) update *services.Update - notifier *notifications.NotificationService + notifier *Notifier loc *Localizer onIconChange func() // onMenuChange drives a full tray relayout: the update row lives in the @@ -35,10 +37,10 @@ type trayUpdater struct { progressWindowOpen bool } -func newTrayUpdater(app *application.App, window *application.WebviewWindow, update *services.Update, notifier *notifications.NotificationService, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { +func newTrayUpdater(app *application.App, showMainAt func(url string), update *services.Update, notifier *Notifier, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { u := &trayUpdater{ app: app, - window: window, + showMainAt: showMainAt, update: update, notifier: notifier, loc: loc, @@ -76,15 +78,15 @@ func (u *trayUpdater) applyLanguage() { u.refreshMenuItem(state) } -// handleClick opens the GitHub releases page when not Enforced, otherwise shows -// the progress page and asks the daemon to start the installer. +// handleClick opens the installer download link when not Enforced, otherwise +// shows the progress page and asks the daemon to start the installer. func (u *trayUpdater) handleClick() { u.mu.Lock() state := u.state u.mu.Unlock() if !state.Enforced { - _ = u.app.Browser.OpenURL(urlGitHubReleases) + _ = u.app.Browser.OpenURL(version.DownloadUrl()) return } @@ -184,14 +186,12 @@ func (u *trayUpdater) sendUpdateNotification(st updater.State) { // openProgressWindow points the main window at the /update progress page and // brings it forward. func (u *trayUpdater) openProgressWindow(version string) { - if u.window == nil { + if u.showMainAt == nil { return } url := "/#/update" if version != "" { - url += "?version=" + version + url += "?version=" + neturl.QueryEscape(version) } - u.window.SetURL(url) - u.window.Show() - u.window.Focus() + u.showMainAt(url) } diff --git a/client/ui/uilogpath.go b/client/ui/uilogpath.go index 96fbb9637..6fa400e01 100644 --- a/client/ui/uilogpath.go +++ b/client/ui/uilogpath.go @@ -8,21 +8,19 @@ import ( log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/configs" "github.com/netbirdio/netbird/client/ui/guilog" ) -// uiLogFileName must stay in sync with the daemon's "gui-client*.log.*" glob -// for rotated siblings (addUILog in client/internal/debug). -const uiLogFileName = "gui-client.log" - // uiLogPath returns the GUI log path with native separators, since the daemon -// opens it directly for debug-bundle collection. +// opens it directly for debug-bundle collection. The file name comes from +// configs.UILogFile so the daemon validates and collects the same name. func uiLogPath() (string, error) { dir, err := os.UserConfigDir() if err != nil { return "", err } - return filepath.Join(dir, "netbird", uiLogFileName), nil + return filepath.Join(dir, "netbird", configs.UILogFile), nil } // newDebugLog builds the GUI debug log, disabled when userSetLogFile is set diff --git a/client/ui/xembed_host_gtk3_linux.go b/client/ui/xembed_host_gtk3_linux.go new file mode 100644 index 000000000..b1f18cc90 --- /dev/null +++ b/client/ui/xembed_host_gtk3_linux.go @@ -0,0 +1,40 @@ +//go:build linux && gtk3 && !(linux && 386) + +package main + +import ( + "errors" + + "github.com/godbus/dbus/v5" +) + +// The legacy GTK3 / WebKit2GTK 4.1 build (-tags gtk3) drops the in-process +// XEmbed StatusNotifierWatcher entirely. The real implementation +// (xembed_host_linux.go + xembed_tray_linux.c) links GTK4 and uses GTK4-only +// popup-menu APIs that have no drop-in GTK3 equivalent, so rather than port the +// C layer we stub the host out on gtk3 builds. The tray still works on every +// desktop that ships its own StatusNotifierWatcher (KDE, GNOME+AppIndicator, +// Cinnamon/xapp, XFCE, …); only the minimal-WM fallback (Fluxbox/OpenBox/i3/ +// dwm/vanilla GNOME) is unavailable on gtk3 packages. See LINUX-TRAY.md. + +// xembedHost is a placeholder so the package compiles on gtk3 builds; the real +// type (with X11/GTK4 state) lives in xembed_host_linux.go. It is never +// instantiated here because xembedTrayAvailable always reports false. +type xembedHost struct{} + +// run satisfies the call in tray_watcher_linux.go; unreachable on gtk3 because +// newXembedHost never returns a non-nil host. +func (*xembedHost) run() {} + +// xembedTrayAvailable always reports false on gtk3 builds, so the watcher probe +// loop in startStatusNotifierWatcher exits immediately and newXembedHost is +// never reached. recenter_linux.go's predicate becomes a harmless no-op too. +func xembedTrayAvailable() bool { + return false +} + +// newXembedHost exists only to satisfy the reference in tray_watcher_linux.go; +// it is unreachable because xembedTrayAvailable returns false on gtk3. +func newXembedHost(conn *dbus.Conn, busName string, objPath dbus.ObjectPath) (*xembedHost, error) { + return nil, errors.New("xembed host unsupported on gtk3 build") +} diff --git a/client/ui/xembed_host_linux.go b/client/ui/xembed_host_linux.go index ba7a5d07c..5551cd02a 100644 --- a/client/ui/xembed_host_linux.go +++ b/client/ui/xembed_host_linux.go @@ -1,4 +1,4 @@ -//go:build linux && !(linux && 386) +//go:build linux && !gtk3 && !(linux && 386) package main diff --git a/client/ui/xembed_tray_linux.c b/client/ui/xembed_tray_linux.c index 07ec74bb2..86da8bf70 100644 --- a/client/ui/xembed_tray_linux.c +++ b/client/ui/xembed_tray_linux.c @@ -1,3 +1,5 @@ +//go:build linux && !gtk3 && !(linux && 386) + #include "xembed_tray_linux.h" #include diff --git a/client/wasm/cmd/main.go b/client/wasm/cmd/main.go index 4683f4033..260a528f0 100644 --- a/client/wasm/cmd/main.go +++ b/client/wasm/cmd/main.go @@ -56,8 +56,7 @@ func startClient(ctx context.Context, nbClient *netbird.Client) error { // parseClientOptions extracts NetBird options from JavaScript object func parseClientOptions(jsOptions js.Value) (netbird.Options, error) { options := netbird.Options{ - DeviceName: "dashboard-client", - LogLevel: defaultLogLevel, + LogLevel: defaultLogLevel, } if jwtToken := jsOptions.Get("jwtToken"); !jwtToken.IsNull() && !jwtToken.IsUndefined() { @@ -87,13 +86,41 @@ func parseClientOptions(jsOptions js.Value) (netbird.Options, error) { options.DeviceName = deviceName.String() } - if disableIPv6 := jsOptions.Get("disableIPv6"); !disableIPv6.IsNull() && !disableIPv6.IsUndefined() { - options.DisableIPv6 = disableIPv6.Bool() + disableIPv6, err := boolOption(jsOptions, "disableIPv6") + if err != nil { + return options, err + } + if disableIPv6 != nil { + options.DisableIPv6 = *disableIPv6 } + // The caller decides whether this client uses lazy connections; left unset it + // defers to the management feature flag. A short-lived, interactive caller + // turns it off so its sessions reach the few peers their grant covers eagerly, + // instead of the first request waiting for the connection to be established. + lazyConnectionEnabled, err := boolOption(jsOptions, "lazyConnectionEnabled") + if err != nil { + return options, err + } + options.LazyConnectionEnabled = lazyConnectionEnabled + return options, nil } +// boolOption reads a boolean option, returning nil when the caller left it out. +// js.Value.Bool panics on any other type, so a wrong type is reported instead. +func boolOption(jsOptions js.Value, name string) (*bool, error) { + v := jsOptions.Get(name) + if v.IsNull() || v.IsUndefined() { + return nil, nil + } + if v.Type() != js.TypeBoolean { + return nil, fmt.Errorf("option %s must be a boolean, got %s", name, v.Type()) + } + b := v.Bool() + return &b, nil +} + // createStartMethod creates the start method for the client func createStartMethod(client *netbird.Client) js.Func { return js.FuncOf(func(this js.Value, args []js.Value) any { diff --git a/client/wasm/cmd/main_test.go b/client/wasm/cmd/main_test.go new file mode 100644 index 000000000..3ec5a8f6a --- /dev/null +++ b/client/wasm/cmd/main_test.go @@ -0,0 +1,64 @@ +//go:build js + +package main + +import ( + "syscall/js" + "testing" +) + +// TestParseClientOptionsBooleans covers the boolean options against the value +// kinds a JS caller can pass: js.Value.Bool panics on anything but a boolean, +// so a wrong type has to be rejected before it reaches the client. +func TestParseClientOptionsBooleans(t *testing.T) { + t.Run("unset leaves the lazy override empty", func(t *testing.T) { + options, err := parseClientOptions(js.Global().Get("Object").New()) + if err != nil { + t.Fatalf("parse options: %v", err) + } + if options.LazyConnectionEnabled != nil { + t.Errorf("lazy override should stay unset, got %v", *options.LazyConnectionEnabled) + } + if options.DisableIPv6 { + t.Error("disableIPv6 should default to false") + } + }) + + t.Run("null defers to the management flag", func(t *testing.T) { + jsOptions := js.Global().Get("Object").New() + jsOptions.Set("lazyConnectionEnabled", js.Null()) + options, err := parseClientOptions(jsOptions) + if err != nil { + t.Fatalf("parse options: %v", err) + } + if options.LazyConnectionEnabled != nil { + t.Errorf("lazy override should stay unset, got %v", *options.LazyConnectionEnabled) + } + }) + + t.Run("booleans are carried through", func(t *testing.T) { + jsOptions := js.Global().Get("Object").New() + jsOptions.Set("lazyConnectionEnabled", false) + jsOptions.Set("disableIPv6", true) + options, err := parseClientOptions(jsOptions) + if err != nil { + t.Fatalf("parse options: %v", err) + } + if options.LazyConnectionEnabled == nil || *options.LazyConnectionEnabled { + t.Errorf("lazy override should be false, got %v", options.LazyConnectionEnabled) + } + if !options.DisableIPv6 { + t.Error("disableIPv6 should be true") + } + }) + + t.Run("a non-boolean is rejected", func(t *testing.T) { + for _, value := range []any{"true", 1, js.Global().Get("Object").New()} { + jsOptions := js.Global().Get("Object").New() + jsOptions.Set("lazyConnectionEnabled", value) + if _, err := parseClientOptions(jsOptions); err == nil { + t.Errorf("value %v should be rejected", value) + } + } + }) +} diff --git a/client/wasm/internal/ssh/client.go b/client/wasm/internal/ssh/client.go index 9cfe65266..28ae95ec0 100644 --- a/client/wasm/internal/ssh/client.go +++ b/client/wasm/internal/ssh/client.go @@ -80,13 +80,12 @@ func (c *Client) Connect(host string, port int, username, jwtToken string, ipVer return fmt.Errorf("dial %s: %w", addr, err) } - sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config) + sshClient, err := nbssh.Handshake(ctx, conn, addr, config) if err != nil { - closeWithLog(conn, "connection after handshake error") - return fmt.Errorf("SSH handshake: %w", err) + return err } - c.sshClient = ssh.NewClient(sshConn, chans, reqs) + c.sshClient = sshClient logrus.Infof("SSH: Connected to %s", addr) return nil @@ -119,57 +118,26 @@ func (c *Client) getAuthMethods(jwtToken string) ([]ssh.AuthMethod, error) { return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil } -// StartSession starts an SSH session with PTY +// StartSession starts an SSH session with PTY. It holds the client lock for +// the whole startup so Close cannot tear the client down mid-setup and the +// new session cannot be installed into an already closed client. func (c *Client) StartSession(cols, rows int) error { + c.mu.Lock() + defer c.mu.Unlock() + if c.sshClient == nil { return fmt.Errorf("SSH client not connected") } - session, err := c.sshClient.NewSession() + pty, err := nbssh.StartPTYSession(c.sshClient, cols, rows) if err != nil { - return fmt.Errorf("create session: %w", err) + return err } - c.mu.Lock() - defer c.mu.Unlock() - c.session = session - - modes := ssh.TerminalModes{ - ssh.ECHO: 1, - ssh.TTY_OP_ISPEED: 14400, - ssh.TTY_OP_OSPEED: 14400, - ssh.VINTR: 3, - ssh.VQUIT: 28, - ssh.VERASE: 127, - } - - if err := session.RequestPty("xterm-256color", rows, cols, modes); err != nil { - closeWithLog(session, "session after PTY error") - return fmt.Errorf("PTY request: %w", err) - } - - c.stdin, err = session.StdinPipe() - if err != nil { - closeWithLog(session, "session after stdin error") - return fmt.Errorf("get stdin: %w", err) - } - - c.stdout, err = session.StdoutPipe() - if err != nil { - closeWithLog(session, "session after stdout error") - return fmt.Errorf("get stdout: %w", err) - } - - c.stderr, err = session.StderrPipe() - if err != nil { - closeWithLog(session, "session after stderr error") - return fmt.Errorf("get stderr: %w", err) - } - - if err := session.Shell(); err != nil { - closeWithLog(session, "session after shell error") - return fmt.Errorf("start shell: %w", err) - } + c.session = pty.Session + c.stdin = pty.Stdin + c.stdout = pty.Stdout + c.stderr = pty.Stderr logrus.Info("SSH: Session started with PTY") return nil diff --git a/combined/cmd/admin.go b/combined/cmd/admin.go new file mode 100644 index 000000000..66fac4ac9 --- /dev/null +++ b/combined/cmd/admin.go @@ -0,0 +1,151 @@ +package cmd + +import ( + "context" + "fmt" + + "github.com/dexidp/dex/storage" + log "github.com/sirupsen/logrus" + "github.com/spf13/cobra" + + "github.com/netbirdio/netbird/formatter/hook" + admincmd "github.com/netbirdio/netbird/management/cmd/admin" + tokencmd "github.com/netbirdio/netbird/management/cmd/token" + nbconfig "github.com/netbirdio/netbird/management/internals/server/config" + "github.com/netbirdio/netbird/management/server/activity" + activitystore "github.com/netbirdio/netbird/management/server/activity/store" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/util" +) + +// newAdminCommands creates the admin command tree with combined-specific resource openers. +func newAdminCommands() *cobra.Command { + return admincmd.NewCommands(admincmd.Openers{ + Resources: withAdminResources, + Store: withAdminStoreOnly, + IDP: withAdminIDPOnly, + }) +} + +func newLegacyTokenCommand() *cobra.Command { + cmd := tokencmd.NewCommands(tokencmd.StoreOpener(withAdminStoreOnly)) + cmd.Deprecated = "use 'admin token' instead" + return cmd +} + +// withAdminResources loads the combined YAML config, initializes stores, and calls fn. +func withAdminResources(cmd *cobra.Command, fn func(ctx context.Context, resources admincmd.Resources) error) error { + return withAdminConfig(cmd, func(ctx context.Context, cfg *CombinedConfig) error { + mgmtConfig, err := adminManagementConfig(cfg) + if err != nil { + return err + } + + managementStore, err := openAdminStore(ctx, cfg) + if err != nil { + return err + } + defer admincmd.CloseStore(ctx, managementStore) + + idpStorage, idpStorageFile, err := admincmd.OpenIDPStorage(mgmtConfig) + if err != nil { + return err + } + defer admincmd.CloseIDPStorage(idpStorage) + + eventStore, esErr := openAdminEventStore(ctx, cfg, mgmtConfig) + if esErr != nil { + _, _ = fmt.Fprintf(cmd.ErrOrStderr(), "Warning: audit events will not be recorded: %v\n", esErr) + } + if eventStore != nil { + defer func() { + if err := eventStore.Close(ctx); err != nil { + log.Debugf("close activity event store: %v", err) + } + }() + } + + return fn(ctx, admincmd.Resources{Store: managementStore, IDPStorage: idpStorage, IDPStorageFile: idpStorageFile, EventStore: eventStore}) + }) +} + +// withAdminStoreOnly opens only the management store for admin subcommands that do not +// need embedded IdP storage. +func withAdminStoreOnly(cmd *cobra.Command, fn func(ctx context.Context, s store.Store) error) error { + return withAdminConfig(cmd, func(ctx context.Context, cfg *CombinedConfig) error { + managementStore, err := openAdminStore(ctx, cfg) + if err != nil { + return err + } + defer admincmd.CloseStore(ctx, managementStore) + + return fn(ctx, managementStore) + }) +} + +func withAdminIDPOnly(cmd *cobra.Command, fn func(ctx context.Context, idpStorage storage.Storage, storageFile string) error) error { + return withAdminConfig(cmd, func(ctx context.Context, cfg *CombinedConfig) error { + mgmtConfig, err := adminManagementConfig(cfg) + if err != nil { + return err + } + idpStorage, idpStorageFile, err := admincmd.OpenIDPStorage(mgmtConfig) + if err != nil { + return err + } + defer admincmd.CloseIDPStorage(idpStorage) + + return fn(ctx, idpStorage, idpStorageFile) + }) +} + +func withAdminConfig(cmd *cobra.Command, fn func(ctx context.Context, cfg *CombinedConfig) error) error { + if err := util.InitLog("error", "console"); err != nil { + return fmt.Errorf("init log: %w", err) + } + + ctx := context.WithValue(cmd.Context(), hook.ExecutionContextKey, hook.SystemSource) //nolint:staticcheck + + cfg, err := LoadConfig(configPath) + if err != nil { + return fmt.Errorf("load config: %w", err) + } + cfg.ApplyAdminDefaults() + applyServerStoreEnv(cfg.Server.Store) + + return fn(ctx, cfg) +} + +func adminManagementConfig(cfg *CombinedConfig) (*nbconfig.Config, error) { + mgmtConfig, err := cfg.ToManagementConfig() + if err != nil { + return nil, fmt.Errorf("create management config: %w", err) + } + return mgmtConfig, nil +} + +func openAdminStore(ctx context.Context, cfg *CombinedConfig) (store.Store, error) { + managementStore, err := store.NewStore(ctx, types.Engine(cfg.Management.Store.Engine), cfg.Management.DataDir, nil, true) + if err != nil { + return nil, fmt.Errorf("create store: %w", err) + } + return managementStore, nil +} + +func openAdminEventStore(ctx context.Context, cfg *CombinedConfig, config *nbconfig.Config) (activity.Store, error) { + if config.DataStoreEncryptionKey == "" { + return nil, fmt.Errorf("data store encryption key is not configured") + } + if err := applyActivityStoreEnv(cfg.Server.ActivityStore); err != nil { + return nil, fmt.Errorf("configure activity event store: %w", err) + } + eventStore, err := activitystore.NewSqlStore(ctx, config.Datadir, config.DataStoreEncryptionKey) + if err != nil { + return nil, fmt.Errorf("open activity event store: %w", err) + } + if eventStore == nil { + return nil, fmt.Errorf("open activity event store: returned nil store") + } + return eventStore, nil +} diff --git a/combined/cmd/admin_config_test.go b/combined/cmd/admin_config_test.go new file mode 100644 index 000000000..ff7045d38 --- /dev/null +++ b/combined/cmd/admin_config_test.go @@ -0,0 +1,47 @@ +package cmd + +import ( + "context" + "os" + "testing" + + "github.com/stretchr/testify/require" + + nbconfig "github.com/netbirdio/netbird/management/internals/server/config" +) + +func TestApplyAdminDefaultsCopiesServerStoreWithoutExposedAddress(t *testing.T) { + cfg := DefaultConfig() + cfg.Server.ExposedAddress = "" + cfg.Server.DataDir = "/srv/netbird" + cfg.Server.Store = StoreConfig{ + Engine: "postgres", + DSN: "postgres://user:pass@example.com/netbird", + } + + cfg.ApplyAdminDefaults() + + require.Equal(t, "/srv/netbird", cfg.Management.DataDir) + require.Equal(t, "postgres", cfg.Management.Store.Engine) + require.Equal(t, cfg.Server.Store.DSN, cfg.Management.Store.DSN) +} + +func TestOpenAdminEventStoreMissingEncryptionKeyReturnsNilInterface(t *testing.T) { + eventStore, err := openAdminEventStore(context.Background(), &CombinedConfig{}, &nbconfig.Config{}) + require.Error(t, err) + require.Contains(t, err.Error(), "encryption key") + require.Nil(t, eventStore) +} + +func TestApplyServerStoreEnv(t *testing.T) { + t.Setenv("NB_STORE_ENGINE_POSTGRES_DSN", "") + t.Setenv("NB_STORE_ENGINE_MYSQL_DSN", "") + t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", "") + + applyServerStoreEnv(StoreConfig{Engine: "postgres", DSN: "postgres-dsn", File: "store.db"}) + require.Equal(t, "postgres-dsn", os.Getenv("NB_STORE_ENGINE_POSTGRES_DSN")) + require.Equal(t, "store.db", os.Getenv("NB_STORE_ENGINE_SQLITE_FILE")) + + applyServerStoreEnv(StoreConfig{Engine: "mysql", DSN: "mysql-dsn"}) + require.Equal(t, "mysql-dsn", os.Getenv("NB_STORE_ENGINE_MYSQL_DSN")) +} diff --git a/combined/cmd/config.go b/combined/cmd/config.go index fcbc60dc9..890c86876 100644 --- a/combined/cmd/config.go +++ b/combined/cmd/config.go @@ -6,8 +6,7 @@ import ( "net" "net/netip" "os" - "path" - "path/filepath" + filePath "path/filepath" "strings" "time" @@ -74,6 +73,16 @@ type ServerConfig struct { ActivityStore StoreConfig `yaml:"activityStore"` AuthStore StoreConfig `yaml:"authStore"` ReverseProxy ReverseProxyConfig `yaml:"reverseProxy"` + + SupportedSyncMessageVersions *int `yaml:"supportedSyncMessageVersions,omitempty"` + PerAccountSupportedSyncMessageVersions map[string]int `yaml:"perAccountSupportedSyncMessageVersions,omitempty"` + + AgentNetwork AgentNetworkConfig `yaml:"agentNetwork"` +} + +// AgentNetworkConfig contains agent-network (LLM gateway) configuration. +type AgentNetworkConfig struct { + PricingDefaultsFile string `yaml:"pricingDefaultsFile"` } // TLSConfig contains TLS/HTTPS settings @@ -300,6 +309,19 @@ func (c *CombinedConfig) ApplySimplifiedDefaults() { c.autoConfigureClientSettings(exposedProto, exposedHost, exposedHostPort, hasExternalStuns, hasExternalRelay, hasExternalSignal) } +// ApplyAdminDefaults applies the management settings needed by admin commands even +// when the full server config is invalid and ApplySimplifiedDefaults cannot run. +func (c *CombinedConfig) ApplyAdminDefaults() { + if c.Management.DataDir == "" || c.Management.DataDir == "/var/lib/netbird/" { + c.Management.DataDir = c.Server.DataDir + } + if c.Management.Store.Engine == "" || c.Management.Store.Engine == "sqlite" { + if c.Server.Store.Engine != "" || c.Server.Store.File != "" || c.Server.Store.DSN != "" { + c.Management.Store = c.Server.Store + } + } +} + // applyRelayDefaults configures the relay service if no external relay is configured. func (c *CombinedConfig) applyRelayDefaults(exposedProto, exposedHostPort string, hasExternalRelay, hasExternalStuns bool) { if hasExternalRelay { @@ -577,11 +599,11 @@ func (c *CombinedConfig) buildEmbeddedIdPConfig(mgmt ManagementConfig) (*idp.Emb return nil, fmt.Errorf("authStore.dsn is required when authStore.engine is postgres") } } else { - authStorageFile = path.Join(mgmt.DataDir, "idp.db") + authStorageFile = filePath.Join(mgmt.DataDir, "idp.db") if c.Server.AuthStore.File != "" { authStorageFile = c.Server.AuthStore.File - if !filepath.IsAbs(authStorageFile) { - authStorageFile = filepath.Join(mgmt.DataDir, authStorageFile) + if !filePath.IsAbs(authStorageFile) { + authStorageFile = filePath.Join(mgmt.DataDir, authStorageFile) } } } @@ -696,16 +718,21 @@ func (c *CombinedConfig) ToManagementConfig() (*nbconfig.Config, error) { httpConfig.AuthCallbackURL = callbackURL + types.ProxyCallbackEndpointFull return &nbconfig.Config{ - Stuns: stuns, - Relay: relayConfig, - Signal: signalConfig, - Datadir: mgmt.DataDir, - DataStoreEncryptionKey: mgmt.Store.EncryptionKey, - HttpConfig: httpConfig, - StoreConfig: storeConfig, - ReverseProxy: reverseProxy, - DisableDefaultPolicy: mgmt.DisableDefaultPolicy, - EmbeddedIdP: embeddedIdP, + Stuns: stuns, + Relay: relayConfig, + Signal: signalConfig, + Datadir: mgmt.DataDir, + DataStoreEncryptionKey: mgmt.Store.EncryptionKey, + HttpConfig: httpConfig, + StoreConfig: storeConfig, + ReverseProxy: reverseProxy, + DisableDefaultPolicy: mgmt.DisableDefaultPolicy, + EmbeddedIdP: embeddedIdP, + HighestSupportedSyncMessageVersion: c.Server.SupportedSyncMessageVersions, + PerAccountHighestSupportedSyncMessageVersion: c.Server.PerAccountSupportedSyncMessageVersions, + AgentNetwork: nbconfig.AgentNetwork{ + PricingDefaultsFile: c.Server.AgentNetwork.PricingDefaultsFile, + }, }, nil } @@ -729,7 +756,7 @@ func ApplyEmbeddedIdPConfig(ctx context.Context, cfg *nbconfig.Config, mgmtPort cfg.EmbeddedIdP.Storage.Type = "sqlite3" } if cfg.EmbeddedIdP.Storage.Config.File == "" && cfg.Datadir != "" { - cfg.EmbeddedIdP.Storage.Config.File = path.Join(cfg.Datadir, "idp.db") + cfg.EmbeddedIdP.Storage.Config.File = filePath.Join(cfg.Datadir, "idp.db") } issuer := cfg.EmbeddedIdP.Issuer diff --git a/combined/cmd/root.go b/combined/cmd/root.go index 2b7956f11..7eac84ce5 100644 --- a/combined/cmd/root.go +++ b/combined/cmd/root.go @@ -10,6 +10,7 @@ import ( "net/http" "os" "os/signal" + "path/filepath" "strconv" "strings" "sync" @@ -24,6 +25,7 @@ import ( "google.golang.org/grpc" "github.com/netbirdio/netbird/encryption" + agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" mgmtServer "github.com/netbirdio/netbird/management/internals/server" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" "github.com/netbirdio/netbird/management/server/telemetry" @@ -31,6 +33,7 @@ import ( relayServer "github.com/netbirdio/netbird/relay/server" "github.com/netbirdio/netbird/relay/server/listener" "github.com/netbirdio/netbird/relay/server/listener/ws" + syncgrpc "github.com/netbirdio/netbird/shared/management/grpc" sharedMetrics "github.com/netbirdio/netbird/shared/metrics" "github.com/netbirdio/netbird/shared/relay/auth" "github.com/netbirdio/netbird/shared/signal/proto" @@ -64,7 +67,8 @@ func init() { rootCmd.PersistentFlags().StringVarP(&configPath, "config", "c", "", "path to YAML configuration file (required)") _ = rootCmd.MarkPersistentFlagRequired("config") - rootCmd.AddCommand(newTokenCommands()) + rootCmd.AddCommand(newAdminCommands()) + rootCmd.AddCommand(newLegacyTokenCommand()) } func RootCmd() *cobra.Command { @@ -122,6 +126,37 @@ func execute(cmd *cobra.Command, _ []string) error { } // initializeConfig loads and validates the configuration, then initializes logging. +func applyServerStoreEnv(storeConfig StoreConfig) { + if dsn := storeConfig.DSN; dsn != "" { + switch strings.ToLower(storeConfig.Engine) { + case "postgres": + os.Setenv("NB_STORE_ENGINE_POSTGRES_DSN", dsn) + case "mysql": + os.Setenv("NB_STORE_ENGINE_MYSQL_DSN", dsn) + } + } + if file := storeConfig.File; file != "" { + os.Setenv("NB_STORE_ENGINE_SQLITE_FILE", file) + } +} + +func applyActivityStoreEnv(storeConfig StoreConfig) error { + if engine := storeConfig.Engine; engine != "" { + engineLower := strings.ToLower(engine) + if engineLower == "postgres" && storeConfig.DSN == "" { + return fmt.Errorf("activityStore.dsn is required when activityStore.engine is postgres") + } + os.Setenv("NB_ACTIVITY_EVENT_STORE_ENGINE", engineLower) + if dsn := storeConfig.DSN; dsn != "" { + os.Setenv("NB_ACTIVITY_EVENT_POSTGRES_DSN", dsn) + } + } + if file := storeConfig.File; file != "" { + os.Setenv("NB_ACTIVITY_EVENT_SQLITE_FILE", file) + } + return nil +} + func initializeConfig() error { var err error config, err = LoadConfig(configPath) @@ -137,30 +172,10 @@ func initializeConfig() error { return fmt.Errorf("failed to initialize log: %w", err) } - if dsn := config.Server.Store.DSN; dsn != "" { - switch strings.ToLower(config.Server.Store.Engine) { - case "postgres": - os.Setenv("NB_STORE_ENGINE_POSTGRES_DSN", dsn) - case "mysql": - os.Setenv("NB_STORE_ENGINE_MYSQL_DSN", dsn) - } - } - if file := config.Server.Store.File; file != "" { - os.Setenv("NB_STORE_ENGINE_SQLITE_FILE", file) - } + applyServerStoreEnv(config.Server.Store) - if engine := config.Server.ActivityStore.Engine; engine != "" { - engineLower := strings.ToLower(engine) - if engineLower == "postgres" && config.Server.ActivityStore.DSN == "" { - return fmt.Errorf("activityStore.dsn is required when activityStore.engine is postgres") - } - os.Setenv("NB_ACTIVITY_EVENT_STORE_ENGINE", engineLower) - if dsn := config.Server.ActivityStore.DSN; dsn != "" { - os.Setenv("NB_ACTIVITY_EVENT_POSTGRES_DSN", dsn) - } - } - if file := config.Server.ActivityStore.File; file != "" { - os.Setenv("NB_ACTIVITY_EVENT_SQLITE_FILE", file) + if err := applyActivityStoreEnv(config.Server.ActivityStore); err != nil { + return err } log.Infof("Starting combined NetBird server") @@ -275,6 +290,11 @@ func (s *serverInstances) createManagementServer(ctx context.Context, cfg *Combi return fmt.Errorf("failed to ensure encryption key: %w", err) } + if err := loadAgentNetworkPricing(ctx, mgmtConfig); err != nil { + cleanupSTUNListeners(s.stunListeners) + return fmt.Errorf("failed to load agent-network pricing defaults: %w", err) + } + LogConfigInfo(mgmtConfig) s.mgmtSrv, err = createManagementServer(cfg, mgmtConfig) @@ -505,6 +525,16 @@ func createManagementServer(cfg *CombinedConfig, mgmtConfig *nbconfig.Config) (m } mgmtPort, _ := strconv.Atoi(portStr) + if err := syncgrpc.ValidateSyncMessageVersion(mgmtConfig.HighestSupportedSyncMessageVersion); err != nil { + return nil, err + } + + for accountId, version := range mgmtConfig.PerAccountHighestSupportedSyncMessageVersion { + if err := syncgrpc.ValidateSyncMessageVersion(&version); err != nil { + return nil, fmt.Errorf("unrecognized sync message version in perAccountSupportedSyncMessageVersions for account %s %w", accountId, err) + } + } + mgmtSrv := newServer( &mgmtServer.Config{ NbConfig: mgmtConfig, @@ -599,6 +629,32 @@ func handleRelayWebSocket(w http.ResponseWriter, r *http.Request, acceptFn func( acceptFn(conn) } +// loadAgentNetworkPricing loads the management-side LLM pricing defaults +// file for the combined server and starts its periodic reloader. An +// explicitly configured PricingDefaultsFile is required to load (a typo +// must fail startup rather than silently bill with built-ins the operator +// believes they replaced); a relative path is resolved against the data +// directory so a bare filename like "pricing.yaml" lands in the datadir +// alongside the store. With no path configured, / +// is probed and may be absent (compiled-in defaults serve). +func loadAgentNetworkPricing(ctx context.Context, mgmtConfig *nbconfig.Config) error { + pricingPath := mgmtConfig.AgentNetwork.PricingDefaultsFile + required := pricingPath != "" + if !required { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(mgmtConfig.Datadir, pricingPath) + } + + log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, required) + if err := agentnetworkpricing.LoadFile(pricingPath, required); err != nil { + return err + } + agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval) + return nil +} + // logConfig prints all configuration parameters for debugging func logConfig(cfg *CombinedConfig) { log.Info("=== Configuration ===") @@ -675,6 +731,25 @@ func logManagementConfig(cfg *CombinedConfig) { log.Infof(" Relay addresses: %v", cfg.Management.Relays.Addresses) log.Infof(" Relay credentials TTL: %s", cfg.Management.Relays.CredentialsTTL) } + + logAgentNetworkConfig(cfg) +} + +func logAgentNetworkConfig(cfg *CombinedConfig) { + log.Info(" Agent Network:") + pricingPath := cfg.Server.AgentNetwork.PricingDefaultsFile + configured := pricingPath != "" + if !configured { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(cfg.Management.DataDir, pricingPath) + } + if configured { + log.Infof(" Pricing defaults file: %s", pricingPath) + } else { + log.Infof(" Pricing defaults file: %s (default, optional)", pricingPath) + } } // logEnvVars logs all NB_ environment variables that are currently set diff --git a/combined/cmd/token.go b/combined/cmd/token.go deleted file mode 100644 index 550480062..000000000 --- a/combined/cmd/token.go +++ /dev/null @@ -1,63 +0,0 @@ -package cmd - -import ( - "context" - "fmt" - "os" - "strings" - - log "github.com/sirupsen/logrus" - "github.com/spf13/cobra" - - "github.com/netbirdio/netbird/formatter/hook" - tokencmd "github.com/netbirdio/netbird/management/cmd/token" - "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" - "github.com/netbirdio/netbird/util" -) - -// newTokenCommands creates the token command tree with combined-specific store opener. -func newTokenCommands() *cobra.Command { - return tokencmd.NewCommands(withTokenStore) -} - -// withTokenStore loads the combined YAML config, initializes the store, and calls fn. -func withTokenStore(cmd *cobra.Command, fn func(ctx context.Context, s store.Store) error) error { - if err := util.InitLog("error", "console"); err != nil { - return fmt.Errorf("init log: %w", err) - } - - ctx := context.WithValue(cmd.Context(), hook.ExecutionContextKey, hook.SystemSource) //nolint:staticcheck - - cfg, err := LoadConfig(configPath) - if err != nil { - return fmt.Errorf("load config: %w", err) - } - - if dsn := cfg.Server.Store.DSN; dsn != "" { - switch strings.ToLower(cfg.Server.Store.Engine) { - case "postgres": - os.Setenv("NB_STORE_ENGINE_POSTGRES_DSN", dsn) - case "mysql": - os.Setenv("NB_STORE_ENGINE_MYSQL_DSN", dsn) - } - } - if file := cfg.Server.Store.File; file != "" { - os.Setenv("NB_STORE_ENGINE_SQLITE_FILE", file) - } - - datadir := cfg.Management.DataDir - engine := types.Engine(cfg.Management.Store.Engine) - - s, err := store.NewStore(ctx, engine, datadir, nil, true) - if err != nil { - return fmt.Errorf("create store: %w", err) - } - defer func() { - if err := s.Close(ctx); err != nil { - log.Debugf("close store: %v", err) - } - }() - - return fn(ctx, s) -} diff --git a/combined/config.yaml.example b/combined/config.yaml.example index 66bc71703..085e4344f 100644 --- a/combined/config.yaml.example +++ b/combined/config.yaml.example @@ -134,3 +134,16 @@ server: # trustedPeers: [] # CIDRs of trusted peer networks (e.g. ["100.64.0.0/10"]) # accessLogRetentionDays: 7 # Days to retain HTTP access logs. 0 (or unset) defaults to 7. Negative values disable cleanup (logs kept indefinitely). # accessLogCleanupIntervalHours: 24 # How often (in hours) to run the access-log cleanup job. 0 (or unset) is treated as "not set" and defaults to 24 hours; cleanup remains enabled. To disable cleanup, set accessLogRetentionDays to a negative value. + + # Agent network (LLM gateway) settings (optional) + # agentNetwork: + # # Path to the YAML file holding the default LLM pricing table. A relative + # # path is resolved against dataDir, so a bare filename like "pricing.yaml" + # # lands in the data directory. When empty, {dataDir}/defaults_llm_pricing.yaml + # # is probed; if no file is present the compiled-in defaults are used. + # # Schema: surface ("openai"/"anthropic"/"bedrock") -> model -> rates in USD + # # per 1k tokens (input_per_1k, output_per_1k, and the optional + # # cached_input_per_1k / cache_read_per_1k / cache_creation_per_1k). The file + # # is re-read periodically (mtime poll). An explicitly configured path that + # # fails to load fails startup; runtime reload errors keep the previous table. + # pricingDefaultsFile: "pricing.yaml" diff --git a/crowdin.yml b/crowdin.yml new file mode 100644 index 000000000..efdbceb8d --- /dev/null +++ b/crowdin.yml @@ -0,0 +1,11 @@ +skip_untranslated_strings: true +skip_untranslated_files: true +import_eq_suggestions: true + +files: + - source: /client/ui/i18n/locales/en/common.json + translation: /client/ui/i18n/locales/%two_letters_code%/common.json + type: chrome + languages_mapping: + two_letters_code: + zh-CN: zh-CN diff --git a/dns/nameserver.go b/dns/nameserver.go index 81c616c50..84e83e2b4 100644 --- a/dns/nameserver.go +++ b/dns/nameserver.go @@ -53,6 +53,7 @@ type NameServerGroup struct { ID string `gorm:"primaryKey"` // AccountID is a reference to Account that this object belongs AccountID string `gorm:"index"` + PublicID string `json:"-"` // Name group name Name string // Description group description diff --git a/docs/agent-networks/01-end-to-end-flows.md b/docs/agent-networks/01-end-to-end-flows.md index 7264f3768..0de6b4c33 100644 --- a/docs/agent-networks/01-end-to-end-flows.md +++ b/docs/agent-networks/01-end-to-end-flows.md @@ -109,13 +109,13 @@ sequenceDiagram Chk->>Inj: continue Inj->>Inj: inject NetBird identity headers per provider config Inj->>Grd: continue - Grd->>Grd: enforce model allowlist + Grd->>Grd: enforce per-provider allowlist (fail-closed backstop) Grd->>Up: forward (over WireGuard) Up-->>Resp: response (JSON or SSE stream) Resp->>Resp: parse usage tokens, completion Note over Resp: capture_completion gates raw
completion capture Resp->>Cost: tokens - Cost->>Cost: lookup pricing.yaml + compute cost + Cost->>Cost: lookup rates from config-delivered
pricing table + compute cost Cost->>Rec: tokens + cost Rec->>MgmtGrpc: RecordLLMUsage(provider, model, prompt_t, completion_t, cost, groups, user) Rec-->>Log: emit access-log entry
(if EnableLogCollection) @@ -135,6 +135,21 @@ sequenceDiagram (`redact_pii = settings.RedactPii`). Phones, emails, credit cards, PII names — see `redact.go` for the full set. See [`modules/31-proxy-middleware-builtin.md`](modules/31-proxy-middleware-builtin.md). +- The model allowlist is enforced in TWO places. `CheckLLMPolicyLimits` + is authoritative: it resolves the policy that governs this + (provider, caller-groups) and denies (`deny_code = llm_policy.model_blocked`) + when no applicable policy permits the model — so an allowlist scoped to + one group/provider never leaks to another, and an un-guardrailed policy + is genuinely unrestricted. `llm_guardrail` is a per-provider fail-closed + backstop: it only carries an allowlist for a provider every authorising + policy restricts, and blocks unknown/undetermined models even when + management is unreachable. Because that backstop allowlist is the UNION + of every restricting policy's models, per-group narrowing lives only in + the authoritative check: during a `CheckLLMPolicyLimits` outage + `llm_limit_check` fails open, so a caller can reach any model in the + provider's union — a group scoped to model A could reach model B if + another group restricts the same provider to B. This is the documented + fail-open trade-off; a future flag may switch it to fail-closed. - SSE streaming requires special handling on the response side; the parser must handle partial chunks without buffering the whole stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.md). diff --git a/docs/agent-networks/modules/21-management-agentnetwork.md b/docs/agent-networks/modules/21-management-agentnetwork.md index b64c1ba20..f91c369f7 100644 --- a/docs/agent-networks/modules/21-management-agentnetwork.md +++ b/docs/agent-networks/modules/21-management-agentnetwork.md @@ -15,6 +15,10 @@ Inside the package: `manager.go` is the CRUD + permissions-gated facade; `synthe | ---- | ---- | | `agentnetwork/manager.go` | Manager interface + CRUD + permission gates + bootstrap-settings + reconcile trigger | | `agentnetwork/synthesizer.go` | Settings/policy → wire-format synthesis; sole writer of the proxy middleware chain | +| `agentnetwork/synthesizer_pricing.go` | `buildCostMeterConfigJSON` — default table + per-provider prices → `cost_meter` config | +| `agentnetwork/pricing/defaults.go` | Default pricing table derived from the catalog + supplementals; `DefaultTable`, `LookupDefault`, wire `Entry` | +| `agentnetwork/pricing/override.go` | `LoadFile`/`StartReloader` for `AgentNetwork.PricingDefaultsFile` (mtime poll, merge over compiled-in base) | +| `agentnetwork/pricing/{exampleyaml,gen}.go` | Generates `defaults_llm_pricing.example.yaml` from the compiled-in table (golden-tested) | | `agentnetwork/policyselect.go` | Per-request policy attribution + account-budget ceiling (min-wins) | | `agentnetwork/reconcile.go` | Per-account synth diff vs in-memory cache → Create/Update/Delete | | `agentnetwork/catalog/catalog.go` | Static provider catalogue (auth headers, identity-injection shapes) | @@ -48,6 +52,8 @@ flowchart TD I --> J[indexProviderGroups: providerID -> sorted source groups] J --> K[buildRouterConfigJSON drops orphan providers] J --> L[buildIdentityInjectConfigJSON per catalog entry] + J --> K2[buildCostMeterConfigJSON: default table + per-provider prices] + K2 --> P H --> M[mergeGuardrails: union allowlist, OR redact] M --> N[applyAccountCollectionControls account toggle = SOLE capture control] N --> O[marshalGuardrailConfig] @@ -60,6 +66,84 @@ flowchart TD R --> T[accountManager.UpdateAccountPeers — fans synth ACLs into network map] ``` +### LLM pricing (management is the sole authority) + +**The proxy carries no price list.** Management synthesizes the entire pricing +table and ships it inside `cost_meter`'s `ConfigJSON`, so a price change reaches +the proxies as an ordinary mapping push — the chain rebuild installs a fresh +table and there is nothing to reload on the proxy side. + +```mermaid +flowchart TD + A[catalog.All — PricingSurfaces x Models] --> B[buildDefaultTable + supplementalDefaults] + B --> C{AgentNetwork.PricingDefaultsFile} + C -- absent --> D[compiled-in table serves] + C -- loaded --> E[LoadFile: merge file entries WHOLE over compiled base] + E --> F[mergedTable atomic.Pointer] + D --> G[DefaultTable] + F --> G + G --> H[buildCostMeterConfigJSON — pricing.defaults] + I[types.Provider.Models operator prices] --> J[normalizePricingModelID
bedrock ARN/region/version, vertex @version] + J --> K[materializeEntry: default entry as base,
operator input/output verbatim,
cache pointers only when non-nil] + K --> L[pricing.providers keyed by provider record ID] + H --> M[cost_meter ConfigJSON] + L --> M + G --> N[GET /catalog — applyDefaultPricing prefills dashboard rows] + O[StartReloader: mtime poll every ReloadInterval 1m] --> E +``` + +**Two tiers, resolved per request on the proxy** (`synthesizer_pricing.go:22-35`): + +- `pricing.defaults` — surface (`openai`/`anthropic`/`bedrock`) → normalized model + id → rates. The **full** default table ships to every account: it is small + (~10 KB) and it is what keeps gateway-style providers (which enumerate no + models, so they claim every model) priced. +- `pricing.providers` — provider **record** id → normalized model id → rates, + matched against the `llm.resolved_provider_id` the router stamps. Entries are + **fully materialized here**, at synth time: `materializeEntry` starts from the + default entry for that model so cache rates the operator didn't state are + inherited, overlays operator `input`/`output` verbatim (**including an explicit + 0**, which prices a self-hosted or internal endpoint as free rather than + silently reverting to list price), and overlays cache-rate **pointers only when + non-nil** — `nil` means "inherit the default", an explicit `0` means "no + discount, bill this bucket at the input rate". The proxy therefore does two map + lookups and no merging. + +Same orphan rule as the router: a provider no enabled policy authorises is +unreachable, so its prices aren't shipped. Model ids are normalized with the +**same** functions the request parser uses (`NormalizeBedrockModel` / +`NormalizeVertexModel`), which is what makes the per-record lookup key compare +equal to the `llm.model` the proxy meters. Post-normalization duplicates resolve +first-occurrence-wins, matching the routing dedup order. + +**`AgentNetwork.PricingDefaultsFile`** (`config.go:190-207`) lets an operator +replace default rates without a rebuild. Schema is `surface → model → rates` +(`input_per_1k`, `output_per_1k`, and optional `cached_input_per_1k` / +`cache_read_per_1k` / `cache_creation_per_1k`). Semantics: + +- A **relative** path resolves against ``, so a bare filename lands + alongside the store. Empty config probes `/defaults_llm_pricing.yaml`. +- An **explicitly configured** path is *required to load*: a typo or malformed + file fails startup, because the operator believes those rates are live. The + conventional probe is optional — an absent file just serves compiled-in + defaults, and the path stays watched in case it appears later. +- File entries **replace** the compiled-in entry for the same (surface, model) + **whole** — they are not field-merged, so an entry must repeat the cache rates + it wants to keep. Everything the file doesn't mention keeps built-in rates. +- Unknown YAML fields are rejected (`KnownFields(true)`) and every rate must be + finite and non-negative — the same constraints the HTTP API enforces on + operator per-provider prices. +- Reload is an mtime poll (`ReloadInterval`, 1 min) and is **lenient at runtime**: + a parse error keeps the previous table, a deleted file reverts to compiled-in + defaults. A mid-edit save can never take pricing down. + +The live table feeds **both** consumers, which is what keeps them consistent: the +synthesizer (what proxies actually bill with) and `GET /api/agent-network/catalog` +via `applyDefaultPricing` (what the dashboard's model-row prices prefill with). +`defaults_llm_pricing.example.yaml` is generated from the compiled-in table +(`go generate ./management/internals/modules/agentnetwork/pricing`) and +golden-tested, so operators start from a file matching the built-in rates exactly. + ### Budget rule resolution (min-wins, group+user bound) ```mermaid @@ -122,9 +206,9 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest | on_request | 1 | `llm_router` | `{"providers":[{id, models[], upstream_*, auth_header_*, allowed_group_ids[]}]}` | **true** | | on_request | 2 | `llm_limit_check` | `{}` | – | | on_request | 3 | `llm_identity_inject` | `{"providers":[{provider_id, header_pair?, json_metadata?, extra_headers?}]}` | **true** | - | on_request | 4 | `llm_guardrail` | `{"model_allowlist"?, "prompt_capture":{enabled,redact_pii}}` | – | + | on_request | 4 | `llm_guardrail` | `{"provider_allowlists"?: {providerID: []model}, "prompt_capture":{enabled,redact_pii}}` | – | | on_response | 5 | `llm_limit_record` | `{}` (runs LAST at runtime) | – | - | on_response | 6 | `cost_meter` | `{}` | – | + | on_response | 6 | `cost_meter` | `{"pricing":{"defaults":{surface:{model:rates}},"providers"?:{providerRecordID:{model:rates}}}}` — rates are `{input_per_1k, output_per_1k, cached_input_per_1k?, cache_read_per_1k?, cache_creation_per_1k?}` | – | | on_response | 7 | `llm_response_parser` | `{"capture_completion": , "redact_pii"?: true}` | – | - **Synthesized service shape** (`synthesizer.go:739`): `Mode=HTTP`, `Private=true`, `Domain=.`, `AccessGroups=unionSourceGroups(enabledPolicies)`, one `TargetTypeCluster` target with `Host=noop.invalid:443` (router rewrites per request), `Options.{DirectUpstream,AgentNetwork}=true`, `DisableAccessLog=!settings.EnableLogCollection`, `CaptureMax{Req,Resp}Bytes=1<<20`, `CaptureContentTypes=["application/json","text/event-stream"]`. @@ -139,6 +223,12 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest - **Orphan providers (no enabled policy authorises them) NEVER reach the router** (`synthesizer.go:351-357`); skipped from `identity_inject` for symmetry. - **Provider creation refuses empty `api_key`** (`manager.go:175`); **deletion refuses while any policy still references it** (`manager.go:265-273`). - **Session keypair stability across provider edits** (`manager.go:226-228`) — server-managed, copied through every `UpdateProvider`, never API-surfaced. +- **Management is the sole pricing authority.** The proxy has no embedded price list, so an account whose `cost_meter` config carries no `pricing` block bills **nothing** (`cost.skipped=unknown_model`, $0) rather than falling back to stale built-ins. The top-level `pricing` wrapper is also the feature-detection signal in both directions: an old proxy ignores it as an unknown field, and a new proxy reads its absence as "old management". +- **Per-provider prices are materialized at synth time, not merged on the proxy** (`synthesizer_pricing.go:114-131`). A per-record entry is always complete, so the proxy's lookup is per-record-then-defaults with no field-level fallback between tiers. +- **An explicit operator price of `0` prices the model as free** — it must not be treated as "unset" and reverted to list price (`synthesizer_pricing.go:49-54`). Only *cache*-rate fields distinguish unset from zero, via `*float64`. +- **Pricing model ids are normalized with the same functions the request parser uses** (`normalizePricingModelID`). If the two ever diverge, per-record prices silently stop matching and every request falls through to surface defaults. +- **The default table's coverage is structural, not curated.** It is derived from the catalog via each provider's `PricingSurfaces`; `TestDefaultTable_CoversEveryCatalogModel` fails on an unpriced catalog model and `TestDefaultTable_NoConflictingContributions` fails if two providers contribute the same (surface, model) at different rates. +- **A pricing-defaults file failure is fatal only at startup, and only for an explicitly configured path.** Runtime reload failures keep the previous table; a deleted file reverts to compiled-in defaults (`pricing/override.go:62-81, 113-148`). ## Things to scrutinize @@ -176,10 +266,12 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest - **Capture-pointer semantics (restated):** non-agent-network callers see no field → legacy nil-default emit, identical to pre-PR. Agent-network targets always carry an explicit `capture_*` value. - **`TestSynthesizeServices_HappyPath` was updated:** request-parser config moved from `{}` to `{"capture_prompt":false}` (`synthesizer_test.go:174`). External snapshot tests against synth output need updating. - **`MergedGuardrails` retains zeroed `TokenLimits`/`Budget`/`Retention`** even though `Policy.Limits` carries the real values now; `llm_limit_check` is the authoritative enforcement. Comment at `synthesizer.go:940-948` calls this out. +- **`cost_meter`'s `pricing` block is version-skew-safe in both directions.** A proxy predating config-delivered pricing ignores the field as unknown JSON (it previously priced from its own embedded table, so it keeps billing — at its own rates, which is the skew to be aware of during a rolling upgrade). A current proxy paired with old management sees no `pricing` block, logs one warning at chain-build time, and records `cost.skipped=unknown_model` — token counting and cap enforcement are unaffected, only the USD annotation goes to $0. ### Performance - **`SynthesizeServices` runs on every controller tick / mutation reconcile.** Cost: 4 store reads + optional per-provider keypair backfill. Sort + index + merge are O(N log N) / O(P × G); dominant cost is JSON marshalling. No nested loops escape these dimensions. +- **The full default pricing table is marshalled into every account's `cost_meter` config on every synth** (~10 KB serialized). This is a deliberate trade: it keeps gateway-style providers priced for every catalog model, and it is the largest single contributor to the synth JSON. `DefaultTable()` itself is a pointer load (or a `sync.Once`-built map) — the cost is the marshal, not the build. - **`reconcile.diffMappings` is O(N + M)** with N=M=1 per account today — effectively constant. - **`SynthesizeServicesForCluster`** (`synthesizer.go:71`) walks every account on a cluster; per-account failures are **swallowed** (`synthesizer.go:91-93`) so a single misconfigured account doesn't drop the cluster. Runs per proxy reconnect. @@ -188,6 +280,7 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest - **Activity codes:** `AgentNetwork{Provider,Policy,Guardrail,BudgetRule}{Created,Updated,Deleted}`; `AgentNetworkSettingsUpdated` with `log_collection/prompt_collection/redact_pii` payload (`manager.go:567-571`). **No activity code for `SelectPolicyForRequest` denies** — surfaced via proxy access log only (likely intentional given volume). - **Deny codes** namespaced: `llm_policy.{token,budget}_cap_exceeded`, `llm_account.{token,budget}_cap_exceeded` (`policyselect.go:18-26`). - **Reconcile failures are logged at warn and swallowed** (`reconcile.go:42-44`). Persistent synth failures (e.g. unknown catalog id) silently keep the proxy out of sync — consider a manager-level synth-health surface if this becomes a support burden. +- **Pricing-file lifecycle logs at info** (load, reload, revert-to-built-ins) and **at warn** for a runtime reload failure; the mtime check itself is `Debugf`. There is no metric on reload failures, so an operator who breaks the file mid-flight keeps billing at the previous table with only a log line to show it (`pricing/override.go:113-148`). ## Test coverage @@ -198,6 +291,9 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest | `synthesizer_guardrail_realstore_test.go` | `PromptCaptureAccountIsSoleControl`; `PromptCaptureFlowsWhenAccountOptsIn`; `AccountRedactWithoutGuardrailRedact`; `NoGuardrail_CaptureOff`. | | `synthesizer_log_collection_realstore_test.go` | `LogCollection{Off_SuppressesAccessLog,On_PermitsAccessLog}` — verifies `DisableAccessLog` propagation through `ToProtoMapping`. | | `synthesizer_parser_redact_realstore_test.go` | **Capture-pointer regression suite:** `ParserConfigsCarryRedactPii`; `ParserConfigsSuppressCaptureWhenLogCollectionOnly` (log=on/prompt=off ⇒ both capture flags false); `ParserConfigsOmitRedactPiiWhenOff`. | +| `synthesizer_pricing_test.go` | `BuildCostMeterConfig_{BedrockModelNormalization,CacheRateNilVsZero,OrphanAndGatewayProviders}` — the per-record tier's three load-bearing rules: keys normalized like the parser's, `nil` cache pointer inherits vs explicit `0` bills at input rate, and orphan / gateway (empty `Models`) providers ship no per-record entry. | +| `pricing/defaults_test.go` | `DefaultTable_{CoversEveryCatalogModel,NoConflictingContributions,AllRatesFiniteNonNegative,PinnedRates}`; `LookupDefault_SurfaceOrder`. Catalog-derived coverage + rate sanity are structural, not curated. | +| `pricing/override_test.go` | `LoadFile_{MergesOverCompiledDefaults,MissingPath,RejectsInvalid}`; `Reload_LifeCycle` (mtime detect, parse error keeps previous, delete reverts to built-ins); `ExampleYAML_InSyncWithBuiltins` golden. | | `policyselect_test.go` | Mock-store: `NoApplicablePolicies`; `AllowWithLowestGroupAttribution`; `LargerPoolWinsAcrossUsageLevels`; `StaysOnLargerPoolAfterPartialDrain`; `FallsThroughToSmallerPoolWhenLargerExhausted`; `TiebreakBy{LargerGroupPool,CreatedAt}`; `DeniesWhenAllExhausted`; `UncappedPolicyAlwaysWinsAgainstCapped`; `DisabledPolicyIgnored`; `StoreErrorPropagates`; `RejectsEmptyAccount`; `SharesGroupCounterAcrossPolicies`; `AntiFallThroughOnLowestGroup`; `BudgetOnlyExhaustionDenies`; `BudgetTighterThanTokenWins`. | | `policyselect_realstore_test.go` | Real-sqlite regression guard: `NoApplicablePolicies`; `AllowAndLowestGroupAttribution`; `LargerPoolWins_FallsThroughWhenExhausted`; `BudgetCapDenies`; `GroupCounterSharedAcrossPolicies`; `DisabledPolicyIgnored`. | | `policyselect_account_realstore_test.go` | Account budget rules: `AccountCeilingBindsEvenWithUncappedPolicy` (min-wins); `AccountGroupCeiling`; `AccountTargetUsersBindsOnlyThatUser`; `AccountRuleRecordsToOwnWindow`. | diff --git a/docs/agent-networks/modules/31-proxy-middleware-builtin.md b/docs/agent-networks/modules/31-proxy-middleware-builtin.md index 904de6424..ad56feb77 100644 --- a/docs/agent-networks/modules/31-proxy-middleware-builtin.md +++ b/docs/agent-networks/modules/31-proxy-middleware-builtin.md @@ -5,7 +5,7 @@ LLM request. The two highest-blast-radius areas are the **capture-pointer semantics** and the **limit_check ⇒ limit_record** record-once invariant. Sibling module: [32-proxy-llm-parsers.md](./32-proxy-llm-parsers.md) — the SDK -adapters + pricing catalog this chain delegates to. +adapters + pricing table and cost formula this chain delegates to. --- @@ -34,7 +34,7 @@ rewrites. | `llm_identity_inject` | OnRequest | `llm.{resolved_provider_id,authorising_groups}`, `Input.{UserEmail,UserID,UserGroups,UserGroupNames}` | none | header strip/inject + optional body rewrite | | `llm_guardrail` | OnRequest | `llm.{model,request_prompt_raw}` | `llm_policy.{decision,reason}`, `llm.request_prompt` | none (model allowlist deny) | | `llm_response_parser` | OnResponse | `llm.provider`, `Input.{RespHeaders,RespBody,Status}` | `llm.{input,output,total,cached_input,cache_creation}_tokens`, `llm.response_completion` | none | -| `cost_meter` | OnResponse | `llm.{provider,model}`, token buckets | `cost.usd_total` or `cost.skipped` | pricing lookup | +| `cost_meter` | OnResponse | `llm.{provider,model,resolved_provider_id}`, token buckets | `cost.usd_{input,cached_input,cache_creation,output,total,cache}` or `cost.skipped` | none (in-memory pricing lookup) | | `llm_limit_record` | OnResponse | `llm.{attribution_group_id,attribution_window_seconds,input_tokens,output_tokens}`, `cost.usd_total` | none | gRPC `RecordLLMUsage` | [all_test.go:26–40](../../../proxy/internal/middleware/builtin/all_test.go) @@ -44,7 +44,7 @@ locks the ID set; adding or removing one is a conscious extension. | File | LOC | Notes | |---|---:|---| -| `builtin.go` | 86 | Registry + `FactoryContext` (ctx, data dir, meter, logger, mgmt client) | +| `builtin.go` | 90 | Registry + `FactoryContext` (ctx, meter, logger, mgmt client) | | `all_test.go` | 41 | Locks the 8-ID registry surface | | `agentnetwork_chain_integration_test.go` | 319 | Live sqlite + real gRPC bufconn; gate→recorder wire path | | `llm_request_parser/*` | 162 / 66 / 356 | Provider detection, body parse, prompt extraction with capture-pointer gating | @@ -53,7 +53,7 @@ locks the ID set; adding or removing one is a conscious extension. | `llm_identity_inject/*` | 440 / 108 / 666 | HeaderPair (LiteLLM) + JSONMetadata (Portkey) + ExtraHeaders | | `llm_guardrail/*` | 176 / 82 / 75 / 219 / 217 | Model allowlist + optional prompt capture with PII redaction | | `llm_response_parser/*` | 258 / 222 / 43 / 433 / 169 / 111 | Buffered + SSE accumulation; AWS event-stream accumulator (`streaming_bedrock.go`) for Bedrock; capture-pointer gates completion emit | -| `cost_meter/*` | 181 / 84 / 439 | Token → USD via `proxy/internal/llm/pricing` | +| `cost_meter/*` | 236 / 98 / 586 | Token → USD via `proxy/internal/llm/pricing`; both pricing tiers arrive in the middleware config | | `llm_limit_record/*` | 144 / 35 / 191 | Post-flight `RecordLLMUsage` (5s, debug-on-error) | ## Per-middleware @@ -168,12 +168,46 @@ token schema. ### cost_meter -Reads `llm.provider` + `llm.model` + token buckets, looks up per-1k rate via -`pricing.Loader`, emits `cost.usd_total` or a closed-set `cost.skipped` -reason (`missing_provider/model/tokens`, `unparseable_tokens`, `zero_tokens`, -`unknown_model`). Loader's hot-reload goroutine is bound to proxy-lifetime -context via `startReloader`. **Key invariant:** provider-shape switch lives -in `pricing.Table.Cost` (sibling doc) — `cost_meter` stays provider-agnostic. +Reads `llm.provider` + `llm.model` + token buckets, looks up the per-1k rates, +and emits the full `cost.usd_*` breakdown (four per-bucket values plus the +`_total` and `_cache` aggregates) or a closed-set `cost.skipped` reason +(`missing_provider/model/tokens`, `unparseable_tokens`, `zero_tokens`, +`unknown_model`). + +**Management owns pricing.** The proxy carries no embedded price list: the whole +table arrives in this middleware's `ConfigJSON` as +`{pricing: {defaults, providers}}`, synthesized by management from the catalog +plus the operator's stored per-provider prices +([factory.go:13–34](../../../proxy/internal/middleware/builtin/cost_meter/factory.go)). +Both tiers are validated by `pricing.NewTable` / `pricing.NewEntries` at +construction, so a non-finite or negative rate fails the chain build. A price +change is an ordinary mapping push — the chain rebuild yields a fresh instance +over a fresh immutable table, so there is no data dir, no pricing file, no +reload goroutine, and nothing to invalidate. + +**Two-tier lookup** +([middleware.go:165–183](../../../proxy/internal/middleware/builtin/cost_meter/middleware.go)): + +1. **Per-provider-record** — the operator's stored price for the route that + actually served the request, keyed by the `llm.resolved_provider_id` that + `llm_router` stamped on the allow path, then by normalized model id. Entries + arrive fully materialized (management folds default cache rates in at synth + time), so there is no merging here. Absent metadata — no router in the chain + — skips this tier. +2. **Surface defaults** — the catalog-derived table keyed by `llm.provider` + (`openai`/`anthropic`/`bedrock`). This is also what prices gateway-style + providers, which enumerate no models and therefore get no per-record entry. + +**Backward compatibility:** a config with no `pricing` block means management +predates config-delivered pricing. The factory logs one warning at build time +and the instance records `cost.skipped=unknown_model` ($0) for every request +rather than falling back to a stale built-in price list +([factory.go:55–60](../../../proxy/internal/middleware/builtin/cost_meter/factory.go)). + +**Key invariant:** the provider-shape switch lives in `pricing.EntryCosts` +(sibling doc) and is selected by the **surface**, not by which tier the entry +came from — `cost_meter` stays provider-agnostic, and a per-record override on +an Anthropic route still bills its cache buckets additively. ### llm_limit_record @@ -244,14 +278,16 @@ no mocks. Tests: `TestChain_AllowPath_StampsAttributionAndRecordsCounter` | `llm_router` | `{providers: [{id, models, upstream_scheme, upstream_host, upstream_path?, auth_header_name, auth_header_value, allowed_group_ids}]}` | | `llm_limit_check` | `{}` — pulls `MgmtClient` from `FactoryContext` | | `llm_identity_inject` | `{providers: [{provider_id, header_pair?|json_metadata?, extra_headers?}]}` | -| `llm_guardrail` | `{model_allowlist: []string, prompt_capture: {enabled, redact_pii}}` | +| `llm_guardrail` | `{provider_allowlists: {providerID: []string}, prompt_capture: {enabled, redact_pii}}` — allowlist keyed by resolved provider id; a provider absent from the map is unrestricted (fail-closed backstop; authoritative per-policy/group check is management's `CheckLLMPolicyLimits`) | | `llm_response_parser` | `{redact_pii?, capture_completion?: *bool}` | -| `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) | +| `cost_meter` | `{pricing: {defaults: {surface: {model: rates}}, providers: {providerRecordID: {model: rates}}}}` — rates are `{input_per_1k, output_per_1k, cached_input_per_1k?, cache_read_per_1k?, cache_creation_per_1k?}`. A missing `pricing` key means "management predates config-delivered pricing": every request records `cost.skipped=unknown_model` | | `llm_limit_record` | `{}` — same pattern as `llm_limit_check` | All factories accept empty / null / `{}` / whitespace as zero-value config; only structurally invalid JSON is rejected so misconfig surfaces at chain -build time. +build time. `cost_meter` adds a semantic check on top of that: a `pricing` +block carrying a negative or non-finite rate fails the build too, rather than +mispricing live traffic. ## Invariants @@ -320,10 +356,11 @@ non-object `metadata` field — header path still attributes, but body-level tag-budget enforcement doesn't run for that request. -**Concurrency.** `cost_meter` shares a `pricing.Loader` via -`atomic.Pointer[Table]`; readers always see a consistent table. Every -middleware is a stateless value receiver. Integration test uses real bufconn -gRPC — race detector is the meaningful bar. +**Concurrency.** `cost_meter`'s two pricing tables are built once from the +middleware config and never mutated, so the lookup path needs no lock or atomic +swap — a price change replaces the whole instance. Every middleware is +otherwise a stateless value receiver. Integration test uses real bufconn gRPC — +race detector is the meaningful bar. **Perf.** Hot path is `lookupKV` linear scan over <10 KVs; `cost_meter.Cost` is O(1); SSE accumulation is single-pass. No map allocation per call. @@ -349,13 +386,13 @@ counter accuracy. | `llm_guardrail/redact_test.go` | 15 | Email, SSN, phone (E.164 + NA), bearer, IPv4; fixture-driven | | `llm_response_parser/middleware_test.go` | 18 | Buffered OAI+Anthro, capture-pointer, redact, truncation | | `llm_response_parser/streaming_test.go` | 7 | OAI usage frame, Anthro message_delta, truncated body best-effort | -| `cost_meter/middleware_test.go` | 17 | Each skip reason, provider-shape, pricing loader integration | +| `cost_meter/middleware_test.go` | 22 | Each skip reason, provider-shape formulas, config-delivered defaults, per-record-beats-defaults + miss-falls-back, per-record uses surface formula, nil-pricing skips everything, invalid-rate rejection | | `llm_limit_record/middleware_test.go` | 7 | Skip-on-no-signal, skip-on-missing-attribution, RPC failure swallowed | ## Cross-references - Sibling: [32-proxy-llm-parsers.md](./32-proxy-llm-parsers.md) — SDK adapters - + SSE framer + pricing loader. + + SSE framer + pricing table and cost formula. - Path-routed providers (Vertex AI + Bedrock), `keyfile::` credential, GCP token minting, `/bedrock` prefix: [50-path-routed-providers.md](./50-path-routed-providers.md). diff --git a/docs/agent-networks/modules/32-proxy-llm-parsers.md b/docs/agent-networks/modules/32-proxy-llm-parsers.md index 0376bc988..52faeaac1 100644 --- a/docs/agent-networks/modules/32-proxy-llm-parsers.md +++ b/docs/agent-networks/modules/32-proxy-llm-parsers.md @@ -9,7 +9,7 @@ pricing table's per-provider cost formula is the highest-leverage place a small bug would silently mis-bill operators. Sibling module: [31-proxy-middleware-builtin.md](./31-proxy-middleware-builtin.md) -— the 8 middlewares that consume this package's parsers + pricing loader. +— the 8 middlewares that consume this package's parsers + pricing table. --- @@ -24,8 +24,9 @@ proxy-framework dependencies: - `openai.go` / `anthropic.go` / `bedrock.go` — per-provider `Parser` impls. - `sse.go` — SSE scanner (`Scanner`, `Event`, `NewScanner`). - `errors.go` — sentinels callers branch on with `errors.Is`. -- `pricing/` — embedded-default + hot-reload override table with - symlink-safe Unix loader (build-tagged stub elsewhere). +- `pricing/` — immutable pricing table + the per-surface cost formula. The + rates themselves come from management inside `cost_meter`'s middleware + config; this package holds no price list and reads no files. - `fixtures/` — captured request/response/stream bodies the tests replay. The package carries zero proxy-framework dependencies so the same parsers can @@ -47,12 +48,9 @@ be reused later by a WASM adapter | `sse_test.go` | 175 | 12 tests; fixture replay + multiline + size limits | | `parser_test.go` | 53 | `Parsers()`, `DetectParser`, provider enum values | | `errors.go` | 31 | 6 sentinels: `Err{Unknown,Unsupported}Provider/Model`, `Err{NotLLM,Malformed}Response`, `ErrStreamingUnsupported`, `ErrMalformedRequest` | -| `pricing/pricing.go` | 421 | `Loader`, `Table`, `Entry`; embedded defaults + atomic swap + mtime reload | -| `pricing/pricing_unix.go` | 69 | `O_NOFOLLOW` + fstat-from-FD + 1 MiB cap | -| `pricing/pricing_other.go` | 21 | Stub returning "not supported on this platform" | -| `pricing/pricing_test.go` | 432 | 21 tests — symlink rejection, reload race, path traversal, oversize | -| `pricing/defaults_pricing.yaml` | 85 | go:embed source of truth | -| `fixtures/*` | 21–59 | OAI chat/responses/stream + Anthro messages/stream + pricing starter | +| `pricing/pricing.go` | 234 | `Table`, `Entry`, `EntryJSON`, `Costs`; `NewTable`/`NewEntries` validation + `EntryCosts` formula. No I/O, no reload, no embedded rates | +| `pricing/pricing_test.go` | 177 | 10 tests — provider-shape formulas, cached clamp, rate fallback, nil-safety, rate validation | +| `fixtures/*` | 21–59 | OAI chat/responses/stream + Anthro messages/stream | ## Request body → parser dispatch @@ -188,9 +186,11 @@ response leg, covering both Bedrock body shapes: `totalTokens`). `firstNonZero` folds the two naming conventions into one `Usage`; when Converse omits `totalTokens` the parser sums the buckets. -`ProviderName()` returns `"bedrock"` — its own `defaults_pricing.yaml` block, -keyed by the **normalised** model id (region prefix + version suffix stripped by -the request parser). `ParseResponse` returns `ErrStreamingUnsupported` for an +`ProviderName()` returns `"bedrock"` — its own pricing surface in the table +management ships, keyed by the **normalised** model id (region prefix + version +suffix stripped by the request parser; management normalises its keys the same +way at synth time so the two compare equal). `ParseResponse` returns +`ErrStreamingUnsupported` for an AWS binary event-stream content-type (`application/vnd.amazon.eventstream`, `isAWSEventStream`) so the caller routes to the streaming accumulator instead. @@ -205,11 +205,34 @@ response body. Streaming accumulators live in the middleware package ([llm_response_parser/streaming.go](../../../proxy/internal/middleware/builtin/llm_response_parser/streaming.go)) but use `llm.NewScanner` so the framing contract stays here. -### Pricing catalog +### Pricing table -`Table.Cost` -([pricing.go:129–174](../../../proxy/internal/llm/pricing/pricing.go)) -is the cost formula — most security-relevant math in this module: +**Management is the sole pricing authority.** The proxy carries no embedded +price list and reads no pricing file: the whole table arrives inside +`cost_meter`'s `ConfigJSON` on the ordinary mapping push, and a price change +is just another push — the chain rebuild constructs a fresh `Table`, so there +is nothing to reload +([pricing.go:1–7](../../../proxy/internal/llm/pricing/pricing.go)). The +management side of the contract (catalog defaults, the operator's stored +per-provider prices, and `AgentNetwork.PricingDefaultsFile`) is covered in the +management-side module guide; `cost_meter`'s wire shape is in +[31-proxy-middleware-builtin.md](./31-proxy-middleware-builtin.md). + +`EntryJSON` +([pricing.go:36–45](../../../proxy/internal/llm/pricing/pricing.go)) is the +management→proxy contract — five USD-per-1k rates under `input_per_1k`, +`output_per_1k`, `cached_input_per_1k`, `cache_read_per_1k`, +`cache_creation_per_1k`. Management's `pricing.Entry` marshals the identical +names, and `EntryJSON`/`Entry` are field-identical so `NewEntries` converts by +direct struct conversion rather than field-by-field copying (a new rate can't +be silently dropped in transit). + +`EntryCosts` +([pricing.go:183–234](../../../proxy/internal/llm/pricing/pricing.go)) +is the cost formula — most security-relevant math in this module. The +**surface** (the `llm.provider` value the request parser stamped) selects the +formula, never the tier the entry came from: a per-provider-record override on +an Anthropic route still bills its cache buckets additively. | Provider | Formula | |---|---| @@ -218,7 +241,7 @@ is the cost formula — most security-relevant math in this module: | default | `inTokens × InputPer1K + outTokens × OutputPer1K` | `bedrock` shares the Anthropic additive-cache formula -([pricing.go:172-174](../../../proxy/internal/llm/pricing/pricing.go)): +([pricing.go:214–229](../../../proxy/internal/llm/pricing/pricing.go)): Anthropic-on-Bedrock reports the same additive cache buckets, while non-Anthropic Bedrock models (Nova, Llama) simply report zero in those buckets so cost reduces to `input + output`. @@ -226,15 +249,12 @@ to `input + output`. Each per-bucket rate falls back to `InputPer1K` when zero — operators opt in to discounts by setting the field. -`Loader` -([pricing.go:212–268](../../../proxy/internal/llm/pricing/pricing.go)) -overlays an optional `pricing.yaml` from data-dir on top of the go:embed -defaults. Atomic pointer swap means readers never observe a partial update. -The mtime-poll reloader (30s default cadence) keeps the previous table on -parse failure so cost annotation never goes blank during a botched edit. - -`defaults_pricing.yaml` is the source of truth for built-in pricing. -Operator overrides only carry the entries they want to change. +`Costs` +([pricing.go:143–163](../../../proxy/internal/llm/pricing/pricing.go)) is the +per-request split. The four per-bucket fields are the base; `TotalUSD` and +`CacheUSD` are **derived** in `newCosts` so the aggregates can never drift from +the breakdown. `InputUSD` is always the non-cached input bucket on both +provider shapes, so input and cached-input never double-count. ## Public contracts @@ -264,29 +284,38 @@ Order matters: `DetectFromURL` ties resolve by registration order. `ProviderBedrock = 3`. Numeric values are persisted in nothing today but treat them as wire-stable — new providers must take fresh numbers. -**`Pricing` lookup** -([pricing.go:129](../../../proxy/internal/llm/pricing/pricing.go)): +**`Pricing` construction + lookup** +([pricing.go:60–130](../../../proxy/internal/llm/pricing/pricing.go)): ```go +func NewEntries(raw map[string]map[string]EntryJSON) (map[string]map[string]Entry, error) +func NewTable(raw map[string]map[string]EntryJSON) (*Table, error) + +func (t *Table) Lookup(provider, model string) (Entry, bool) func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (float64, bool) +func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) +func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, cacheCreation int64) Costs ``` -Nil-safe: `t.Cost` on a nil receiver returns `(0, false)` -([pricing.go:130–132](../../../proxy/internal/llm/pricing/pricing.go)). -`ok=false` means provider or model is absent from the loaded table; the caller -emits `cost.skipped=unknown_model`. +`NewTable` is the surface-keyed defaults table; `NewEntries` returns the raw +two-level map `cost_meter` uses for the per-provider-record tier (it looks up an +`Entry` directly and calls `EntryCosts`, so it needs no `Table` wrapper). Both +reject any non-finite or negative rate, so a corrupt config fails the chain +build rather than mispricing silently. Nil input yields an empty, +never-matching table. + +Nil-safe: `t.Cost`/`t.Lookup` on a nil receiver returns `ok=false` +([pricing.go:96–99](../../../proxy/internal/llm/pricing/pricing.go)). +`ok=false` means the surface or model is absent from the table management sent; +the caller emits `cost.skipped=unknown_model`. ## Invariants -1. **Cross-platform pricing build.** `pricing_unix.go` carries the only - functional `loadPricing` (uses `syscall.O_NOFOLLOW` and `f.Stat()` on an - open descriptor — both Unix-only). `pricing_other.go` is a build-tag - fallback that returns `"not supported on this platform"` - ([pricing_other.go:14–16](../../../proxy/internal/llm/pricing/pricing_other.go)). - The proxy is Linux-only in production today; a Windows port needs an - equivalent path-as-handle implementation. Reviewers building on Windows - should expect this surface to return an error at startup if an override - file is configured. +1. **The pricing package is pure and platform-independent.** No file I/O, no + `//go:embed`, no goroutines, no build tags — the rates arrive as config, so + there is nothing platform-specific left to port. Anything reintroducing a + read-from-disk path here re-splits pricing authority between management and + the proxy, which is exactly what this design removed. 2. **SSE scanner handles partial chunks.** A buffered prefix that doesn't end in `\n\n` still yields its accumulated event before `io.EOF` @@ -298,38 +327,45 @@ emits `cost.skipped=unknown_model`. usage rather than aborting ([streaming.go:68–73, 144–150](../../../proxy/internal/middleware/builtin/llm_response_parser/streaming.go)). -3. **`defaults_pricing.yaml` is the source of truth.** Compiled into the - binary via `//go:embed` - ([pricing.go:29–30](../../../proxy/internal/llm/pricing/pricing.go)). - `DefaultTable()` parses once and panics on parse failure - ([pricing.go:42–49](../../../proxy/internal/llm/pricing/pricing.go)) - — by design: a broken embedded YAML must not ship to production. +3. **Management is the only source of rates.** `Table` has no constructor that + invents prices: the only way in is `NewTable`/`NewEntries` over the wire map + management sent. A missing or empty `pricing` block therefore means *no + prices at all* (`cost_meter` records `cost.skipped=unknown_model`, $0) — + never a stale built-in fallback that would silently bill list price. -4. **Loader path validation.** `resolveMiddlewareDataPath` - ([pricing.go:370–394](../../../proxy/internal/llm/pricing/pricing.go)) - rejects absolute paths, traversal segments, and basenames that fail - `basenameRegex = ^[a-zA-Z0-9._-]+$`. The resolved path must remain - inside `baseDir` even after `filepath.Clean`. Tests: - `TestNewLoader_PathValidation`, `TestNewLoader_PathValidation_Extended`, - `TestNewLoader_SymlinkOutsideBaseDirRejected`, `TestNewLoader_SymlinkRejected`. +4. **Tables are immutable once built.** `Table.entries` is written only in + `NewEntries` and never mutated afterwards, and `cost_meter`'s `perRecord` + map is likewise build-time-only + ([pricing.go:47–52](../../../proxy/internal/llm/pricing/pricing.go)). This + is what makes the no-reload design safe: a price change arrives as a mapping + push that builds a new middleware instance over a new table, so concurrent + readers can't observe a half-updated price list and no atomic swap or lock + is needed on the hot path. -5. **Unix loader symlink safety.** `O_NOFOLLOW` on open, `f.Stat()` on the - open descriptor (never re-stat by path), `info.Mode().IsRegular()` check, - `io.LimitReader(f, maxPricingBytes+1)` with a final size assertion - ([pricing_unix.go:25–57](../../../proxy/internal/llm/pricing/pricing_unix.go)). - A mid-read symlink swap is detected because the fstat is on the original - fd. Test: `TestNewLoader_RejectsOversizedFile_FixesM4`. +5. **Rate validation happens at chain-build time, not per request.** + `NewEntries` rejects negative, NaN, and ±Inf rates field by field + ([pricing.go:60–83](../../../proxy/internal/llm/pricing/pricing.go)), naming + the offending surface/model/field in the error. Management enforces the same + constraints at its API boundary and in its YAML parser, so this is + defense-in-depth — but it means a corrupt push fails loudly at build instead + of producing negative costs on live traffic. Test: + `TestNewTable_ValidatesRates`. -6. **`yaml.NewDecoder(...).KnownFields(true)`** - ([pricing.go:397–398](../../../proxy/internal/llm/pricing/pricing.go)) - rejects YAML files that carry fields not in the schema. A typo in an - operator override file fails loud instead of silently zeroing rates. +6. **New rates must be added to `Entry`, `EntryJSON`, *and* management's + `pricing.Entry` together.** `NewEntries` converts by direct struct + conversion `Entry(e)` + ([pricing.go:76–78](../../../proxy/internal/llm/pricing/pricing.go)), which + only compiles while the two structs stay field-identical — so the proxy half + is compiler-enforced. The management half is not: a rate added there but not + here unmarshals into nothing and prices that bucket at `InputPer1K`. ## Things to scrutinise -**Correctness.** Verify OpenAI cached-prompt clamp at -[pricing.go:147–149](../../../proxy/internal/llm/pricing/pricing.go) -short-circuits before subtraction. `Anthropic.TotalTokens` sums all four +**Correctness.** Verify the OpenAI cached-prompt clamp at +[pricing.go:203–206](../../../proxy/internal/llm/pricing/pricing.go) +short-circuits before subtraction. Negative token counts are clamped to zero up +front ([pricing.go:186–197](../../../proxy/internal/llm/pricing/pricing.go)) so +no formula can yield a negative cost. `Anthropic.TotalTokens` sums all four buckets (in + out + cache_read + cache_creation) — downstream dashboards need to know this differs from `input + output`. `OpenAIParser.ExtractPrompt` falls through `messages → input → prompt`; a @@ -338,22 +374,27 @@ noting). **Security.** `Scanner.maxLine = 1 MiB`; a 2 MiB single-line `data:` event errors from `Scanner.Next` and both accumulators stop with partial usage. -Pricing file 1 MiB cap is orders of magnitude larger than realistic. Confirm -new schema additions are mirrored in both `pricingFile` and `Entry`; -`KnownFields(true)` will reject silently-typo'd operator overrides -otherwise. +Pricing is no longer file-backed, so the loader's path-traversal / symlink / +oversize surface is gone entirely — the config channel (an authenticated +mapping push from management) is now the only way rates enter the proxy, and +`NewEntries` is the validation boundary on it. A new rate added to management's +`pricing.Entry` but not to `EntryJSON` here is the remaining silent-mispricing +path (see invariant 6). -**Concurrency.** `Loader.table` is `atomic.Pointer[Table]`; readers never -block or see a torn table. `Loader.Reload` is one goroutine, cancelled via -context (`TestLoader_ReloadBackgroundLoopCancellation`). `DefaultTable()` -uses `sync.Once`. Per-call `Scanner` instances mean no shared state across -concurrent response-parser calls. +**Concurrency.** Nothing in this package is shared mutable state: tables are +built once and never written again, so `cost_meter`'s hot path is lock-free by +construction rather than by atomic swap. Per-call `Scanner` instances mean no +shared state across concurrent response-parser calls. -**Perf.** `Table.Cost` is two map lookups + multiplications, O(1). -`Scanner.Next` is one `ReadString('\n')` per line. Pricing reload poll 30s. +**Perf.** `Table.Cost` is two map lookups + multiplications, O(1); the +per-provider-record tier adds at most one more lookup. `Scanner.Next` is one +`ReadString('\n')` per line. No background goroutines and no per-request +allocation of pricing state. -**Observability.** Reload failures count via `metric.Int64Counter` keyed -`plugin`; warning log rate-limited at 5 min so a broken file doesn't flood. +**Observability.** A config carrying no `pricing` block logs one warning at +chain-build time (`cost_meter` factory) and then records +`cost.skipped=unknown_model` per request, so an old-management deployment is +visible in both logs and the access log rather than quietly reporting $0. Parser errors return sentinels — middleware uses `errors.Is` to map to the right `cost.skipped` reason. @@ -365,7 +406,7 @@ right `cost.skipped` reason. | `openai_test.go` | 11 | Chat Completions + Responses API + legacy `prompt`; cached-tokens subset for both naming conventions; fixture replays | | `anthropic_test.go` | 7 | Messages + legacy `/v1/complete`; streaming REJECTED on `ParseResponse` (must use scanner); fixture replays | | `sse_test.go` | 12 | Fixture replay both providers; multiline `data:`; CRLF; comment skip; trailing-event-without-blank-line; oversize rejection | -| `pricing/pricing_test.go` | 21 | Provider-shape switch; cached-rate fallback; cached-clamp; symlink rejection (target outside basedir + symlink to file); path validation matrix; oversize rejection; reload-keeps-previous-on-parse-error; mtime change detection; goroutine cancellation | +| `pricing/pricing_test.go` | 10 | Provider-shape switch (surface selects the formula); cached-rate + cache-read/creation fallback to `InputPer1K`; cached-clamp; negative-token clamp; nil-receiver safety; rate validation (negative / NaN / Inf rejected); nil + empty table | **Fixtures** ([proxy/internal/llm/fixtures/](../../../proxy/internal/llm/fixtures/)): `openai_chat_completion.json` (chat.completions with usage), @@ -373,14 +414,15 @@ right `cost.skipped` reason. `openai_stream.txt` (3 deltas + usage + `[DONE]`), `anthropic_messages.json` (Messages API non-streaming), `anthropic_stream.txt` (full 7-event sequence: message_start → -content_block_{start,delta×2,stop} → message_delta (usage) → message_stop), -`pricing.yaml` (realistic-pricing starter for operator overrides). +content_block_{start,delta×2,stop} → message_delta (usage) → message_stop). +No pricing fixture: the table is config-delivered, so pricing tests construct +it in-process from a wire-shape map. ## Cross-references - Sibling: [31-proxy-middleware-builtin.md](./31-proxy-middleware-builtin.md) — the chain that calls `llm.Parsers()`, `llm.ParserByName`, - `llm.NewScanner`, `pricing.NewLoader`. + `llm.NewScanner`, `pricing.NewTable` / `pricing.NewEntries`. - Path-routed providers (Vertex AI + Bedrock), credential syntax, and the Bedrock AWS event-stream accumulator: [50-path-routed-providers.md](./50-path-routed-providers.md). diff --git a/docs/agent-networks/modules/33-proxy-runtime.md b/docs/agent-networks/modules/33-proxy-runtime.md index f553473f8..54046b614 100644 --- a/docs/agent-networks/modules/33-proxy-runtime.md +++ b/docs/agent-networks/modules/33-proxy-runtime.md @@ -1,7 +1,7 @@ # proxy/runtime — translate + serve + log > **Risk level:** High — every config push from management is translated here, and the chain runs on every HTTP request to a synth target. -> **Backward-compat impact:** Additive at the wire (`PathTargetOptions.middlewares`, `agent_network`, `disable_access_log`, capture caps) and on the proxy `Server` struct (`MiddlewareDataDir`, `MiddlewareCaptureBudgetBytes`). Non-agent-network targets stay on the no-middleware fast path. +> **Backward-compat impact:** Additive at the wire (`PathTargetOptions.middlewares`, `agent_network`, `disable_access_log`, capture caps) and on the proxy `Server` struct (`MiddlewareCaptureBudgetBytes`). Non-agent-network targets stay on the no-middleware fast path. Middleware config is entirely wire-delivered — no proxy-side data dir is involved, including for LLM pricing, which management ships inside `cost_meter`'s config. ## Module boundary @@ -114,8 +114,7 @@ At **request time** the access-log middleware stamps `CapturedData`; the auth ch ## Public contracts touched -- `proxy.Server.MiddlewareDataDir` (string) — base dir for file-backed middleware config (server.go:238-241). -- `proxy.Server.MiddlewareCaptureBudgetBytes` (int64) — process-wide capture cap; defaults to 256 MiB (server.go:248-250). +- `proxy.Server.MiddlewareCaptureBudgetBytes` (int64) — process-wide capture cap; defaults to 256 MiB (server.go:249-253). There is no `MiddlewareDataDir`: no built-in middleware reads config from disk, so `builtin.FactoryContext` carries only the proxy-lifetime context, meter, logger, and management client. - `proxy/internal/proxy.WithMiddlewareManager(*middleware.Manager) Option` — new option on `NewReverseProxy`; nil keeps the fast path (reverseproxy.go:48-56). - `proxy/internal/proxy.PathTarget` adds `Middlewares`, `CaptureConfig`, `AgentNetwork`, `DisableAccessLog` (servicemapping.go:27-51), all zero-default. - `proxy/internal/proxy.CapturedData` adds `agentNetwork`, `suppressAccessLog`, `userGroupNames` behind `sync.RWMutex`; slices deep-copied (context.go:47-66, 183-258). diff --git a/docs/agent-networks/modules/50-path-routed-providers.md b/docs/agent-networks/modules/50-path-routed-providers.md index b7cda3a97..08c976c5f 100644 --- a/docs/agent-networks/modules/50-path-routed-providers.md +++ b/docs/agent-networks/modules/50-path-routed-providers.md @@ -87,9 +87,9 @@ strips the `@version` suffix from the model, and maps the publisher to a parser surface via `vertexPublisherVendor`: - `anthropic` → `llm.provider="anthropic"` → metered through the Anthropic - parser, priced under the **`anthropic`** block in `defaults_pricing.yaml` - (the parser emits the standard Anthropic provider label, so Vertex Claude - reuses first-party Anthropic prices). + parser, priced under the **`anthropic`** surface of the pricing table + management ships (the parser emits the standard Anthropic provider label, so + Vertex Claude reuses first-party Anthropic prices). - `openai` → `llm.provider="openai"` (reserved; not in the catalog lineup today). - anything else (notably `google` / Gemini) → empty vendor → **no parser**. @@ -104,8 +104,9 @@ is omitted from the catalog. > Caveat: cross-region inference profiles in `eu` / `apac` carry a ~10% price > premium that the base per-token rates do **not** model — cost annotations for -> those regions read low. Operators who need exact regional billing override -> the affected entries in `pricing.yaml`. +> those regions read low. Operators who need exact regional billing set the +> affected models' prices on the provider record, or replace the default entries +> via management's `AgentNetwork.PricingDefaultsFile`. ## AWS Bedrock (`bedrock_api`) @@ -211,15 +212,19 @@ so a model-listing call can't be rewritten onto an upstream that would 404 it. ## Catalog ↔ pricing cross-check Catalog prices and context windows are cross-checked against LiteLLM's -`model_prices_and_context_window.json`. The proxy's embedded -`defaults_pricing.yaml` covers **every metered first-party model** the catalog -enumerates — guarded by -`TestDefaultTable_FirstPartyModelCoverage` -([pricing/defaults_coverage_test.go](../../../proxy/internal/llm/pricing/defaults_coverage_test.go)), -which fails if a catalog model has no embedded price. Bedrock entries are keyed -by the **normalised** id the request parser emits (region prefix + version -suffix stripped). Vertex Claude carries no Bedrock-style prefix, so it prices -straight off the `anthropic` block. +`model_prices_and_context_window.json`. The **catalog is the source of default +prices**: management's `pricing.DefaultTable` folds every catalog provider's +models into the surfaces that provider declares (`PricingSurfaces`), so coverage +is structural rather than maintained in a parallel file +([pricing/defaults.go](../../../management/internals/modules/agentnetwork/pricing/defaults.go)). +`TestDefaultTable_CoversEveryCatalogModel` fails if a catalog model ends up +unpriced, and `TestDefaultTable_NoConflictingContributions` fails if two +providers contribute the same (surface, model) at different rates. Bedrock +entries are keyed by the **normalised** id the request parser emits (region +prefix + version suffix stripped) — management applies the same normalisation to +per-provider prices at synth time, so the two keys compare equal. Vertex Claude +carries no Bedrock-style prefix, so it prices straight off the `anthropic` +surface. ## Things to scrutinise @@ -232,16 +237,17 @@ operator-misconfigured Vertex provider and unmetered Gemini traffic; verify publishers). **Correctness.** `normalizeBedrockModel` is the join between the wire id and the -pricing key — a model that normalises to something not in `defaults_pricing.yaml` -meters at `cost.skipped=unknown_model` rather than failing the request. The +pricing key — a model that normalises to something absent from the shipped +pricing table meters at `cost.skipped=unknown_model` rather than failing the +request. The `/bedrock` prefix strip must run on both the parser side (so the model is extracted) and the router side (so the upstream path is native); a regression in either silently breaks the other. **Metering caveats.** eu/apac cross-region Bedrock + Vertex profiles carry a -~10% premium not modelled by base pricing — flagged in both the catalog comment -and `defaults_pricing.yaml`. Operators needing exact regional billing override -the relevant entries. +~10% premium not modelled by base pricing — flagged in the catalog comment. +Operators needing exact regional billing set per-provider prices on the model +rows (or replace the default entries via `AgentNetwork.PricingDefaultsFile`). ## Cross-references diff --git a/docs/io.netbird.client.plist b/docs/io.netbird.client.plist index 800ecead1..fe10b5b63 100644 --- a/docs/io.netbird.client.plist +++ b/docs/io.netbird.client.plist @@ -66,6 +66,9 @@ disableAutoConnect + disableAutostart + + disableClientRoutes diff --git a/docs/netbird-macos.mobileconfig b/docs/netbird-macos.mobileconfig index 9bf616094..8216dd55d 100644 --- a/docs/netbird-macos.mobileconfig +++ b/docs/netbird-macos.mobileconfig @@ -103,6 +103,8 @@