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..727bff45a 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -475,6 +475,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 +814,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 +826,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 +856,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 +913,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/test-infrastructure-files.yml b/.github/workflows/test-infrastructure-files.yml index 965d8aa5d..729214d9e 100644 --- a/.github/workflows/test-infrastructure-files.yml +++ b/.github/workflows/test-infrastructure-files.yml @@ -257,6 +257,15 @@ jobs: with: persist-credentials: false + - name: Verify fresh-install session cookie key hardening + run: | + 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: Verify Dex retirement notice run: | if infrastructure_files/getting-started-with-dex.sh >stdout.txt 2>stderr.txt; then diff --git a/.goreleaser_ui.yaml b/.goreleaser_ui.yaml index 1157e6379..ca5148823 100644 --- a/.goreleaser_ui.yaml +++ b/.goreleaser_ui.yaml @@ -93,7 +93,9 @@ 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 - maintainer: Netbird description: Netbird client UI. @@ -114,7 +116,9 @@ 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) rpm: signature: diff --git a/.goreleaser_ui_gtk3.yaml b/.goreleaser_ui_gtk3.yaml new file mode 100644 index 000000000..d0a7de6e2 --- /dev/null +++ b/.goreleaser_ui_gtk3.yaml @@ -0,0 +1,131 @@ +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: + # Same package_name as the GTK4 packages -- the two are mutually-exclusive + # alternatives served from separate repo paths (see uploads below); a given + # distro points at exactly one of them. The file names must still differ: + # the Debian pool is shared storage keyed by file name, so a default-named + # gtk3 .deb would overwrite the stable one. + - 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 + file_name_template: "{{ .PackageName }}-gtk3_{{ .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 + dependencies: + - netbird (>= 0.75.0) + - libgtk-3-0 + - libwebkit2gtk-4.1-0 + + - 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 + file_name_template: "{{ .PackageName }}-gtk3_{{ .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 + dependencies: + - netbird >= 0.75.0 + - (gtk3 or libgtk-3-0) + - (webkit2gtk4.1 or libwebkit2gtk-4_1-0) + + 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: + # The gtk3 packages reuse the netbird-ui package name, so they live in + # dedicated repo paths (deb distribution `gtk3`, yum path `yum-gtk3`) that + # legacy distros point their repo config at. + - name: debian-gtk3 + skip: "{{ .Env.SKIP_PUBLISH }}" + ids: + - netbird_ui_deb_gtk3 + mode: archive + target: https://pkgs.wiretrustee.com/debian/pool/{{ .ArtifactName }};deb.distribution=gtk3;deb.component=main;deb.architecture={{ if .Arm }}armhf{{ else }}{{ .Arch }}{{ end }};deb.package= + username: dev@wiretrustee.com + method: PUT + + - name: yum-gtk3 + skip: "{{ .Env.SKIP_PUBLISH }}" + ids: + - netbird_ui_rpm_gtk3 + mode: archive + target: https://pkgs.wiretrustee.com/yum-gtk3/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }} + username: dev@wiretrustee.com + method: PUT diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..95b02a91d --- /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 3b8017788..db5097a48 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -66,11 +66,41 @@ 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) - [Directory structure](#directory-structure) diff --git a/client/android/client.go b/client/android/client.go index 501d7f77c..154bd8484 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -57,6 +57,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()) } @@ -76,6 +82,8 @@ type Client struct { 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 @@ -96,11 +104,12 @@ type Client struct { 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 } @@ -110,6 +119,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() @@ -162,7 +181,7 @@ 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 @@ -170,7 +189,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid // 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) + 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() @@ -211,7 +230,7 @@ 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) + c.setState(cfg, cacheDir, cfgFile, connectClient) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } @@ -240,6 +259,24 @@ 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 +} + // 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) { diff --git a/client/android/login.go b/client/android/login.go index a9422cdbf..3f367b97f 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,17 +178,42 @@ 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 } +// loginHintSetter is implemented by both concrete flows (PKCE and device code) +// but absent from the OAuthFlow interface, hence the assertion below — the same +// way internal/auth wires it in authenticateWithPKCEFlow. +type loginHintSetter interface { + SetLoginHint(hint string) +} + func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) { oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV) if err != nil { return nil, fmt.Errorf("failed to get OAuth flow: %v", err) } + // An empty hint is deliberate, not a fallback: a fresh or logged-out profile + // leaves the choice to the IdP, which is how accounts get switched. + if a.cfgPath != "" { + if hint := readProfileEmail(a.cfgPath); hint != "" { + if setter, ok := oAuthFlow.(loginHintSetter); ok { + setter.SetLoginHint(hint) + } + } + } + flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO()) if err != nil { return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err) diff --git a/client/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/profile_manager.go b/client/android/profile_manager.go index 9a051137c..3197124d7 100644 --- a/client/android/profile_manager.go +++ b/client/android/profile_manager.go @@ -13,18 +13,17 @@ 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 or was logged out. See profile_state.go. + Email string IsActive bool } @@ -101,6 +100,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 +123,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,6 +200,11 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return fmt.Errorf("failed to save config: %w", err) } + // Not fatal: a stale hint costs an account switch, not the logout itself. + if err := removeProfileEmail(configPath); err != nil { + log.Warnf("failed to clear stored account email for profile %s: %v", id, err) + } + log.Infof("logged out from profile: %s", id) return nil } diff --git a/client/android/profile_state.go b/client/android/profile_state.go new file mode 100644 index 000000000..3f0a09701 --- /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 logout: while the +// email is on disk it goes out as a login_hint, which would steer the next login +// straight back into the account just logged out of. 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..623e16c3b --- /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 logout, got %q", got) + } + + // Logout 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/session.go b/client/android/session.go index 961d52528..d5da09c93 100644 --- a/client/android/session.go +++ b/client/android/session.go @@ -278,7 +278,7 @@ func (c *Client) endExtend() { } func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isAndroidTV bool) error { - cfg, _, cc := c.stateSnapshot() + cfg, cfgPath, cc := c.authSnapshot() if cfg == nil || cc == nil { return fmt.Errorf("engine is not running") } @@ -293,7 +293,10 @@ func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isA } defer authClient.Close() - a := &Auth{ctx: ctx, config: cfg} + // 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) diff --git a/client/internal/connect.go b/client/internal/connect.go index 87126b222..ceb39419e 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -113,11 +113,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, 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/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/server.go b/client/internal/dns/server.go index f79454457..3af912792 100644 --- a/client/internal/dns/server.go +++ b/client/internal/dns/server.go @@ -252,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 @@ -602,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 { @@ -686,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) 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/engine.go b/client/internal/engine.go index 617892e43..f4f47992f 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -572,12 +572,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) } @@ -595,10 +590,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, @@ -2102,42 +2095,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.SyncMessageVersion, - 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 { @@ -2172,7 +2129,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() { @@ -2185,7 +2142,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 @@ -2197,7 +2154,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, 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/profilemanager/state.go b/client/internal/profilemanager/state.go index 9e9577796..fcd1c384c 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,18 +82,7 @@ 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 diff --git a/client/internal/routemanager/dnsinterceptor/handler.go b/client/internal/routemanager/dnsinterceptor/handler.go index f92300bfd..d20f4b944 100644 --- a/client/internal/routemanager/dnsinterceptor/handler.go +++ b/client/internal/routemanager/dnsinterceptor/handler.go @@ -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/manager.go b/client/internal/routemanager/manager.go index 2ab7e2a85..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" @@ -62,7 +61,7 @@ type Manager interface { 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 @@ -76,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 @@ -149,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) { @@ -464,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) @@ -500,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 @@ -700,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 diff --git a/client/internal/routemanager/mock.go b/client/internal/routemanager/mock.go index cf761091d..2a8398b95 100644 --- a/client/internal/routemanager/mock.go +++ b/client/internal/routemanager/mock.go @@ -30,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 } 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 c91a76551..d663dd471 100644 --- a/client/internal/routemanager/notifier/notifier_ios.go +++ b/client/internal/routemanager/notifier/notifier_ios.go @@ -29,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 } 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/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/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index 37d3e5d99..2d5460d03 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -158,13 +158,19 @@ 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) {} diff --git a/client/ios/NetBirdSDK/login.go b/client/ios/NetBirdSDK/login.go index 99486839b..6cba0c411 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 := "" 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/server.go b/client/server/server.go index aaab5cc02..01778b8e0 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -135,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 { @@ -370,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 { @@ -623,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 } @@ -684,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 } @@ -839,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 } @@ -1769,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 { 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/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/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/modules/session/SessionExpirationDialog.tsx b/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx index 10e71babb..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"; @@ -131,6 +131,21 @@ export default function SessionExpirationDialog() { } }, [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]); + const logout = useCallback(async () => { if (busy) return; setBusy(true); @@ -185,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/i18n/locales/de/common.json b/client/ui/i18n/locales/de/common.json index 5e91e8d88..d02589591 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" }, diff --git a/client/ui/i18n/locales/en/common.json b/client/ui/i18n/locales/en/common.json index b668146e8..9769e772f 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." diff --git a/client/ui/i18n/locales/es/common.json b/client/ui/i18n/locales/es/common.json index c036e4f75..3420b612b 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" }, diff --git a/client/ui/i18n/locales/fr/common.json b/client/ui/i18n/locales/fr/common.json index c6b91fb25..a83f85c12 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" }, diff --git a/client/ui/i18n/locales/hu/common.json b/client/ui/i18n/locales/hu/common.json index dd5a1af6c..b291f7a01 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" }, diff --git a/client/ui/i18n/locales/it/common.json b/client/ui/i18n/locales/it/common.json index 7a2eb610c..a68a8b32b 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" }, diff --git a/client/ui/i18n/locales/ja/common.json b/client/ui/i18n/locales/ja/common.json index 326c825bf..10cf7598d 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": "表示言語" }, diff --git a/client/ui/i18n/locales/pt/common.json b/client/ui/i18n/locales/pt/common.json index 37b02d5a8..ef1bfd372 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" }, diff --git a/client/ui/i18n/locales/ru/common.json b/client/ui/i18n/locales/ru/common.json index b9ae59df2..a876387f4 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": "Язык интерфейса" }, diff --git a/client/ui/i18n/locales/zh-CN/common.json b/client/ui/i18n/locales/zh-CN/common.json index 2141a770d..542b2b045 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": "显示语言" }, diff --git a/client/ui/main.go b/client/ui/main.go index 782e76d1c..5f740f5ec 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 @@ -102,7 +101,7 @@ func main() { 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 @@ -181,6 +180,7 @@ func main() { WindowManager: windowManager, Session: authSession, Localizer: localizer, + Preferences: prefStore, }) listenForShowSignal(context.Background(), tray) @@ -210,7 +210,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) 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/connection.go b/client/ui/services/connection.go index fae7ddd23..1069f8754 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() } @@ -122,6 +136,7 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err UserCode: resp.GetUserCode(), VerificationURI: resp.GetVerificationURI(), VerificationURIComplete: resp.GetVerificationURIComplete(), + ProfileID: profileID, }, nil } @@ -242,6 +257,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 Logout below). + // 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/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..5a9a0e68d 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,31 @@ 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 (same split as + // Connection.Logout). 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/tray.go b/client/ui/tray.go index 3050d159a..148dd50b3 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" ) @@ -44,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 { @@ -461,10 +463,12 @@ func (t *Tray) handleQuit() { t.profileMu.Unlock() t.svc.DaemonFeed.CancelProfileSwitch() - 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) + 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() } 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 d1117b57b..5b2629419 100644 --- a/client/ui/tray_notify.go +++ b/client/ui/tray_notify.go @@ -44,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 885fdb348..f25419894 100644 --- a/client/ui/tray_session.go +++ b/client/ui/tray_session.go @@ -27,11 +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() } @@ -308,11 +307,7 @@ func (t *Tray) openSessionExtendFlow() { } seconds := int(time.Until(deadline).Seconds()) if seconds <= 0 { - if t.window != nil { - t.window.SetURL("/#/login") - t.window.Show() - t.window.Focus() - } + t.app.Event.Emit(services.EventTriggerLogin) return } if t.svc.WindowManager == nil { diff --git a/client/ui/tray_update.go b/client/ui/tray_update.go index 1a377dfa3..27037eccb 100644 --- a/client/ui/tray_update.go +++ b/client/ui/tray_update.go @@ -21,7 +21,7 @@ type trayUpdater struct { app *application.App window *application.WebviewWindow update *services.Update - notifier *notifications.NotificationService + notifier *Notifier loc *Localizer onIconChange func() // onMenuChange drives a full tray relayout: the update row lives in the @@ -36,7 +36,7 @@ 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, window *application.WebviewWindow, update *services.Update, notifier *Notifier, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { u := &trayUpdater{ app: app, window: window, 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/e2e/agentnetwork/management_test.go b/e2e/agentnetwork/management_test.go index cfd03f63c..6962f9796 100644 --- a/e2e/agentnetwork/management_test.go +++ b/e2e/agentnetwork/management_test.go @@ -160,8 +160,19 @@ func TestSettingsRoundTrip(t *testing.T) { assert.Equal(t, before.Cluster, flipped.Cluster, "cluster must be immutable across updates") assert.Equal(t, before.Subdomain, flipped.Subdomain, "subdomain must be immutable across updates") + // A cluster different from the pinned one must be rejected; echoing the + // pinned one back is valid. + _, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr("attacker.cluster.invalid"), + EnableLogCollection: before.EnableLogCollection, + EnablePromptCollection: before.EnablePromptCollection, + RedactPii: before.RedactPii, + }) + requireClientError(t, err) + // Restore the original toggles. _, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr(before.Cluster), EnableLogCollection: before.EnableLogCollection, EnablePromptCollection: before.EnablePromptCollection, RedactPii: before.RedactPii, diff --git a/e2e/agentnetwork/settings_bootstrap_test.go b/e2e/agentnetwork/settings_bootstrap_test.go new file mode 100644 index 000000000..ea56f7064 --- /dev/null +++ b/e2e/agentnetwork/settings_bootstrap_test.go @@ -0,0 +1,114 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// harnessStartFresh boots a dedicated combined server with its own fresh +// account and registers its teardown on t. +func harnessStartFresh(ctx context.Context, t *testing.T) (*harness.Combined, error) { + t.Helper() + fresh, err := harness.StartCombined(ctx) + if err != nil { + return nil, err + } + t.Cleanup(func() { _ = fresh.Terminate(context.Background()) }) + if _, err := fresh.Bootstrap(ctx); err != nil { + return nil, err + } + return fresh, nil +} + +// TestSettingsBootstrapViaPut covers the settings-first bootstrap path on an +// account that has never been bootstrapped: the GET reads as the defaults +// with an empty cluster/subdomain/endpoint, a PUT without a cluster has +// nothing to pin and fails, and a PUT carrying a cluster creates the row and +// pins it immutably. The shared srv cannot provide that starting state (any +// provider-creating test bootstraps it, and test order is deliberately not +// relied on), so this boots a dedicated combined server — the image is +// already built and cached by TestMain's StartCombined, so the extra cost is +// one container start. +func TestSettingsBootstrapViaPut(t *testing.T) { + ctx := context.Background() + + fresh, err := harnessStartFresh(ctx, t) + require.NoError(t, err, "start dedicated combined server") + + // Before agent-network bootstrap the settings read as the defaults, not + // as an error and not as a null body. + before, err := fresh.GetSettings(ctx) + require.NoError(t, err, "get settings on a fresh account must succeed") + assert.Empty(t, before.Cluster, "cluster must be empty before bootstrap") + assert.Empty(t, before.Subdomain, "subdomain must be empty before bootstrap") + assert.Empty(t, before.Endpoint, "endpoint must be empty before bootstrap, not a bare dot") + assert.True(t, before.EnableLogCollection, "defaults must show log collection on, matching bootstrap") + assert.False(t, before.EnablePromptCollection, "defaults must show prompt collection off") + + // A PUT without a cluster has nothing to pin the account to. + _, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + EnableLogCollection: true, + }) + requireClientError(t, err) + + // A PUT carrying a cluster bootstraps the account and applies the + // mutable fields from the same request. Every toggle is set away from + // its bootstrap default so each assertion can actually fail. + const cluster = "e2e.bootstrap.netbird.selfhosted" + bootstrapped, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr(cluster), + EnableLogCollection: false, + EnablePromptCollection: true, + RedactPii: true, + }) + require.NoError(t, err, "bootstrap settings via PUT must succeed") + assert.Equal(t, cluster, bootstrapped.Cluster, "cluster must be pinned from the request") + require.NotEmpty(t, bootstrapped.Subdomain, "subdomain must be assigned at bootstrap") + assert.Equal(t, bootstrapped.Subdomain+"."+cluster, bootstrapped.Endpoint, "endpoint must combine subdomain and cluster") + assert.False(t, bootstrapped.EnableLogCollection, "log collection from the bootstrap request must override the default") + assert.True(t, bootstrapped.EnablePromptCollection, "prompt collection from the bootstrap request must apply") + assert.True(t, bootstrapped.RedactPii, "redact toggle from the bootstrap request must apply") + + // The row is persisted: an independent read agrees on every field. + after, err := fresh.GetSettings(ctx) + require.NoError(t, err, "get settings after bootstrap must succeed") + assert.Equal(t, bootstrapped.Endpoint, after.Endpoint, "bootstrap must persist across reads") + assert.Equal(t, bootstrapped.EnableLogCollection, after.EnableLogCollection, "log collection must persist") + assert.Equal(t, bootstrapped.EnablePromptCollection, after.EnablePromptCollection, "prompt collection must persist") + assert.Equal(t, bootstrapped.RedactPii, after.RedactPii, "redact toggle must persist") + + // Once bootstrapped, later updates may omit the cluster entirely. + persisted, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + EnableLogCollection: true, + EnablePromptCollection: false, + RedactPii: true, + }) + require.NoError(t, err, "post-bootstrap update without cluster must succeed") + assert.Equal(t, cluster, persisted.Cluster, "omitted cluster must keep the pinned value") + assert.True(t, persisted.EnableLogCollection, "post-bootstrap toggle must apply") + assert.False(t, persisted.EnablePromptCollection, "post-bootstrap toggle must apply") + + // The cluster is immutable: a different value is rejected rather than + // silently ignored, and the rejected update must not disturb anything. + _, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr("other.cluster.invalid"), + EnableLogCollection: false, + }) + requireClientError(t, err) + + final, err := fresh.GetSettings(ctx) + require.NoError(t, err, "get settings after the rejected cluster change must succeed") + assert.Equal(t, persisted.Cluster, final.Cluster, "rejected update must not change the cluster") + assert.Equal(t, persisted.Endpoint, final.Endpoint, "rejected update must not change the endpoint") + assert.Equal(t, persisted.EnableLogCollection, final.EnableLogCollection, "rejected update must not apply its toggles") + assert.Equal(t, persisted.EnablePromptCollection, final.EnablePromptCollection, "rejected update must not apply its toggles") + assert.Equal(t, persisted.RedactPii, final.RedactPii, "rejected update must not apply its toggles") +} diff --git a/go.mod b/go.mod index 26a194958..94b2eb059 100644 --- a/go.mod +++ b/go.mod @@ -114,7 +114,7 @@ require ( github.com/ti-mo/conntrack v0.5.1 github.com/ti-mo/netfilter v0.5.2 github.com/vmihailenco/msgpack/v5 v5.4.1 - github.com/wailsapp/wails/v3 v3.0.0-alpha2.117 + github.com/wailsapp/wails/v3 v3.0.0-beta.3 github.com/yusufpapurcu/wmi v1.2.4 github.com/zcalusic/sysinfo v1.1.3 go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.67.0 @@ -339,3 +339,5 @@ replace github.com/dexidp/dex => github.com/netbirdio/dex v0.244.1-0.20260716205 replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1 replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0 + +replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4 diff --git a/go.sum b/go.sum index 58e30a580..91ce14073 100644 --- a/go.sum +++ b/go.sum @@ -490,6 +490,8 @@ github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9ax github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4 h1:UKztc3QjWvzU5DZk+uYaOWN0x62NSe/pkxuPvzqZIy4= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4/go.mod h1:BzATbK71VFikMMMCo434wAi0QcaI03P+xeaWgDvQvjw= github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a h1:3CWK+yTvRKOcC0Q8VCTGy4l60TEb27CQVS7LkMxwjmw= github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw= github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A= @@ -660,8 +662,6 @@ github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IU github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok= github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g= github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds= -github.com/wailsapp/wails/v3 v3.0.0-alpha2.117 h1:udyjqPG3AIgkod5QDR/WblCkpV8R86BFPSrsWxSyt5Y= -github.com/wailsapp/wails/v3 v3.0.0-alpha2.117/go.mod h1:74WH2FScMsgucZvHHvv7eOefDXCm/CjuIxqhhZgPhKg= github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU= github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA= github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM= diff --git a/infrastructure_files/getting-started-enterprise.sh b/infrastructure_files/getting-started-enterprise.sh index 31f3f1c13..7418cb8e8 100755 --- a/infrastructure_files/getting-started-enterprise.sh +++ b/infrastructure_files/getting-started-enterprise.sh @@ -11,6 +11,10 @@ SED_STRIP_PADDING='s/=//g' NETBIRD_EULA_URL="https://netbird.io/self-hosted-EULA" +# Static IP for Traefik inside the compose bridge network. The management +# server trusts X-Forwarded-* headers from this address only. +TRAEFIK_IP="172.30.0.10" + check_docker_compose() { if command -v docker-compose &> /dev/null; then echo "docker-compose" @@ -80,7 +84,7 @@ read_nb_domain() { if ! check_domain_resolves "$value"; then echo "" > /dev/stderr echo "Warning: '$value' does not resolve via DNS from this host." > /dev/stderr - echo "Caddy will not be able to issue TLS certificates until it does." > /dev/stderr + echo "Traefik will not be able to issue TLS certificates until it does." > /dev/stderr local confirm="" echo -n "Continue anyway? [y/N]: " > /dev/stderr read -r confirm < /dev/tty @@ -92,6 +96,23 @@ read_nb_domain() { echo "$value" } +read_letsencrypt_email() { + if [[ -n "${NETBIRD_LETSENCRYPT_EMAIL:-}" ]]; then + echo "$NETBIRD_LETSENCRYPT_EMAIL" + return + fi + local value="" + echo "Enter your email for Let's Encrypt certificate notifications." > /dev/stderr + echo -n "Email address: " > /dev/stderr + read -r value < /dev/tty + if [[ -z "$value" ]]; then + echo "Email is required for Let's Encrypt." > /dev/stderr + read_letsencrypt_email + return + fi + echo "$value" +} + read_required() { local prompt="$1" local value="" @@ -204,11 +225,11 @@ init_environment() { check_openssl DOCKER_COMPOSE_COMMAND=$(check_docker_compose) - if [[ -f .env ]] || [[ -f docker-compose.yml ]] || [[ -f config.yaml ]] || [[ -f Caddyfile ]]; then + if [[ -f .env ]] || [[ -f docker-compose.yml ]] || [[ -f config.yaml ]]; then echo "Generated files already exist in $(pwd)." echo "If you want to reinitialize the environment, please remove them first:" echo " $DOCKER_COMPOSE_COMMAND down --volumes # removes all containers and volumes" - echo " rm -f .env docker-compose.yml Caddyfile config.yaml" + echo " rm -f .env docker-compose.yml config.yaml" echo "Be aware this will remove all data from the database." exit 1 fi @@ -230,6 +251,9 @@ init_environment() { echo "" NETBIRD_DOMAIN=$(read_nb_domain) + echo "" + NETBIRD_LETSENCRYPT_EMAIL=$(read_letsencrypt_email) + echo "" NETBIRD_LICENSE_KEY=$(read_secret "Enter license key (input hidden)") @@ -238,6 +262,7 @@ init_environment() { POSTGRES_DB="netbird" POSTGRES_PASSWORD=$(rand_secret) NETBIRD_ENCRYPTION_KEY=$(rand_b64_key) + NETBIRD_SESSION_COOKIE_ENCRYPTION_KEY=$(rand_b64_key) NETBIRD_RELAY_AUTH_SECRET=$(rand_secret) POSTGRES_DSN="host=postgres user=${POSTGRES_USER} password=${POSTGRES_PASSWORD} dbname=${POSTGRES_DB} port=5432 sslmode=disable TimeZone=UTC" @@ -247,6 +272,7 @@ init_environment() { echo "Selected:" echo " Traffic flow: ${NETBIRD_TRAFFIC_FLOW}" echo " Domain: ${NETBIRD_DOMAIN}" + echo " ACME email: ${NETBIRD_LETSENCRYPT_EMAIL}" echo "" echo "Rendering files into $(pwd) ..." install -m 600 /dev/null .env @@ -256,7 +282,6 @@ init_environment() { if [[ -z "${NETBIRD_LICENSE_SERVER_BASE_URL:-}" ]]; then sed -i.bak '/NETBIRD_LICENSE_SERVER_BASE_URL/d' docker-compose.yml && rm -f docker-compose.yml.bak fi - render_caddyfile > Caddyfile install -m 600 /dev/null config.yaml render_config_yaml >> config.yaml @@ -283,7 +308,7 @@ init_environment() { echo "All configuration and secrets are stored (mode 600) in $(pwd)/.env" echo "" echo "Tail logs:" - echo " cd $(pwd) && $DOCKER_COMPOSE_COMMAND logs -f netbird-server caddy" + echo " cd $(pwd) && $DOCKER_COMPOSE_COMMAND logs -f netbird-server traefik" } # ------------------------------------------------------------------ @@ -306,6 +331,11 @@ NETBIRD_TRAFFIC_FLOW_ENABLED=${NETBIRD_TRAFFIC_FLOW} # Domain NETBIRD_DOMAIN=${NETBIRD_DOMAIN} +# Reverse proxy (Traefik) +NETBIRD_LETSENCRYPT_EMAIL=${NETBIRD_LETSENCRYPT_EMAIL} +NETBIRD_TRAEFIK_TAG=${NETBIRD_TRAEFIK_TAG:-v3.6} +NETBIRD_TRAEFIK_IP=${TRAEFIK_IP} + # Image tags. Default to "latest" NETBIRD_DASHBOARD_TAG=${NETBIRD_DASHBOARD_TAG:-latest} NETBIRD_SERVER_TAG=${NETBIRD_SERVER_TAG:-latest} @@ -378,26 +408,78 @@ EOF render_compose_common() { cat <<'EOF' - caddy: + # Reverse proxy with automatic TLS via Let's Encrypt. Routes are declared as + # labels on the services below and picked up through the Docker provider. + traefik: <<: *default - image: caddy:2 - container_name: netbird-caddy - networks: [netbird] - environment: - - CADDY_SECURE_DOMAIN=${NETBIRD_DOMAIN} + image: traefik:${NETBIRD_TRAEFIK_TAG} + container_name: netbird-traefik + networks: + netbird: + ipv4_address: ${NETBIRD_TRAEFIK_IP} + command: + # Logging + - "--log.level=INFO" + - "--accesslog=true" + # Docker provider + - "--providers.docker=true" + - "--providers.docker.exposedbydefault=false" + - "--providers.docker.network=netbird" + # Entrypoints + - "--entrypoints.web.address=:80" + - "--entrypoints.websecure.address=:443" + - "--entrypoints.websecure.allowACMEByPass=true" + # readTimeout bounds the whole request, and gRPC streams / relay WebSockets + # never end one; idleTimeout would close the keep-alive connection they + # are reused over. Entrypoint-wide is the only scope Traefik offers here. + # writeTimeout is left alone: it already defaults to 0. + - "--entrypoints.websecure.transport.respondingTimeouts.readTimeout=0" + - "--entrypoints.websecure.transport.respondingTimeouts.idleTimeout=0" + # HTTP to HTTPS redirect + - "--entrypoints.web.http.redirections.entrypoint.to=websecure" + - "--entrypoints.web.http.redirections.entrypoint.scheme=https" + # Let's Encrypt ACME + - "--certificatesresolvers.letsencrypt.acme.email=${NETBIRD_LETSENCRYPT_EMAIL}" + - "--certificatesresolvers.letsencrypt.acme.storage=/letsencrypt/acme.json" + - "--certificatesresolvers.letsencrypt.acme.tlschallenge=true" ports: - '443:443' - - '443:443/udp' - '80:80' volumes: - - netbird_caddy_data:/data - - ./Caddyfile:/etc/caddy/Caddyfile + - /var/run/docker.sock:/var/run/docker.sock:ro + - netbird_traefik_letsencrypt:/letsencrypt + labels: + - traefik.enable=true + # Shared security headers, referenced by every NetBird router below. A + # label-declared middleware only exists while its container runs, so this + # lives on Traefik itself: declaring it on an app container would drop + # every router referencing it whenever that container restarts. + - traefik.http.middlewares.nb-security.headers.stsSeconds=3600 + - traefik.http.middlewares.nb-security.headers.stsIncludeSubdomains=true + - traefik.http.middlewares.nb-security.headers.contentTypeNosniff=true + - traefik.http.middlewares.nb-security.headers.browserXssFilter=true + - traefik.http.middlewares.nb-security.headers.referrerPolicy=strict-origin-when-cross-origin + - traefik.http.middlewares.nb-security.headers.customResponseHeaders.X-Frame-Options=SAMEORIGIN + # Empty value strips the header. Only the dashboard's nginx sets one; the + # server emits none. Do not quote it — "" would send a literal Server: "". + - traefik.http.middlewares.nb-security.headers.customResponseHeaders.Server= dashboard: <<: *default image: ghcr.io/netbirdio/dashboard-cloud:${NETBIRD_DASHBOARD_TAG} container_name: netbird-dashboard networks: [netbird] + labels: + - traefik.enable=true + # Dashboard catch-all: lowest priority so every route below wins + - traefik.http.routers.netbird-dashboard.rule=Host(`${NETBIRD_DOMAIN}`) + - traefik.http.routers.netbird-dashboard.entrypoints=websecure + - traefik.http.routers.netbird-dashboard.tls=true + - traefik.http.routers.netbird-dashboard.tls.certresolver=letsencrypt + - traefik.http.routers.netbird-dashboard.middlewares=nb-security@docker + - traefik.http.routers.netbird-dashboard.service=dashboard + - traefik.http.routers.netbird-dashboard.priority=1 + - traefik.http.services.dashboard.loadbalancer.server.port=80 environment: - NETBIRD_MGMT_API_ENDPOINT=https://${NETBIRD_DOMAIN} - NETBIRD_MGMT_GRPC_API_ENDPOINT=https://${NETBIRD_DOMAIN} @@ -435,6 +517,28 @@ render_compose_server() { - netbird_data:/var/lib/netbird - ./config.yaml:/etc/netbird/config.yaml command: ["--config", "/etc/netbird/config.yaml"] + labels: + - traefik.enable=true + # Signal + Management gRPC (needs an h2c backend for HTTP/2 cleartext) + - traefik.http.routers.netbird-grpc.rule=Host(`${NETBIRD_DOMAIN}`) && (PathPrefix(`/signalexchange.SignalExchange/`) || PathPrefix(`/management.ManagementService/`) || PathPrefix(`/management.ProxyService/`)) + - traefik.http.routers.netbird-grpc.entrypoints=websecure + - traefik.http.routers.netbird-grpc.tls=true + - traefik.http.routers.netbird-grpc.tls.certresolver=letsencrypt + - traefik.http.routers.netbird-grpc.middlewares=nb-security@docker + - traefik.http.routers.netbird-grpc.service=netbird-server-h2c + - traefik.http.routers.netbird-grpc.priority=100 + # Relay WebSocket, management API, and the embedded IdP + - traefik.http.routers.netbird-backend.rule=Host(`${NETBIRD_DOMAIN}`) && (PathPrefix(`/relay`) || PathPrefix(`/ws-proxy/`) || PathPrefix(`/api`) || PathPrefix(`/oauth2`)) + - traefik.http.routers.netbird-backend.entrypoints=websecure + - traefik.http.routers.netbird-backend.tls=true + - traefik.http.routers.netbird-backend.tls.certresolver=letsencrypt + - traefik.http.routers.netbird-backend.middlewares=nb-security@docker + - traefik.http.routers.netbird-backend.service=netbird-server + - traefik.http.routers.netbird-backend.priority=100 + # Services + - traefik.http.services.netbird-server.loadbalancer.server.port=80 + - traefik.http.services.netbird-server-h2c.loadbalancer.server.port=80 + - traefik.http.services.netbird-server-h2c.loadbalancer.server.scheme=h2c environment: - NB_LICENSE_KEY=${NETBIRD_LICENSE_KEY} - NETBIRD_LICENSE_SERVER_BASE_URL=${NETBIRD_LICENSE_SERVER_BASE_URL} @@ -497,6 +601,18 @@ render_compose_flow() { - NB_FLOW_NATS_ENDPOINTS=nats://nats:4222 - NB_FLOW_NATS_STREAM=traffic-events - NB_FLOW_AUTH_SECRET=${NETBIRD_RELAY_AUTH_SECRET} + labels: + - traefik.enable=true + # Flow receiver gRPC (h2c backend) + - traefik.http.routers.netbird-flow.rule=Host(`${NETBIRD_DOMAIN}`) && PathPrefix(`/flow.FlowService/`) + - traefik.http.routers.netbird-flow.entrypoints=websecure + - traefik.http.routers.netbird-flow.tls=true + - traefik.http.routers.netbird-flow.tls.certresolver=letsencrypt + - traefik.http.routers.netbird-flow.middlewares=nb-security@docker + - traefik.http.routers.netbird-flow.service=netbird-flow-h2c + - traefik.http.routers.netbird-flow.priority=100 + - traefik.http.services.netbird-flow-h2c.loadbalancer.server.port=80 + - traefik.http.services.netbird-flow-h2c.loadbalancer.server.scheme=h2c EOF } @@ -536,61 +652,16 @@ EOF fi cat <<'EOF' netbird_postgres: - netbird_caddy_data: + netbird_traefik_letsencrypt: networks: netbird: -EOF -} - -render_caddyfile() { - cat <<'EOF' -{ - servers :80,:443 { - protocols h1 h2c h2 h3 - } -} - -(security_headers) { - header * { - Strict-Transport-Security "max-age=3600; includeSubDomains; preload" - X-Content-Type-Options "nosniff" - X-Frame-Options "SAMEORIGIN" - X-XSS-Protection "1; mode=block" - -Server - Referrer-Policy strict-origin-when-cross-origin - } -} - -:80 { - redir https://{$CADDY_SECURE_DOMAIN}{uri} permanent -} - -{$CADDY_SECURE_DOMAIN}:443 { - import security_headers - # Signal (gRPC over h2c) - reverse_proxy /signalexchange.SignalExchange/* h2c://netbird-server:80 - # Management (gRPC over h2c + HTTP) - reverse_proxy /management.ManagementService/* h2c://netbird-server:80 - reverse_proxy /api/* netbird-server:80 - reverse_proxy /ws-proxy/* netbird-server:80 - # Embedded IdP (OAuth2 endpoints served by netbird server) - reverse_proxy /oauth2/* netbird-server:80 - # Relay (WebSocket multiplexed on the same port) - reverse_proxy /relay* netbird-server:80 -EOF - - if [[ "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then - cat <<'EOF' - # Flow receiver (gRPC over h2c) - reverse_proxy /flow.FlowService/* h2c://receiver:80 -EOF - fi - - cat <<'EOF' - # Dashboard - reverse_proxy /* dashboard:80 -} + name: netbird + driver: bridge + ipam: + config: + - subnet: 172.30.0.0/24 + gateway: 172.30.0.1 EOF } @@ -609,7 +680,7 @@ server: logLevel: "info" logFile: "console" - # TLS is terminated by Caddy in front; leave this block empty. + # TLS is terminated by Traefik in front; leave this block empty. tls: certFile: "" keyFile: "" @@ -626,12 +697,23 @@ server: issuer: "https://${NETBIRD_DOMAIN}/oauth2" localAuthDisabled: false signKeyRefreshEnabled: false + sessionCookieEncryptionKey: "${NETBIRD_SESSION_COOKIE_ENCRYPTION_KEY}" dashboardRedirectURIs: - "https://${NETBIRD_DOMAIN}/nb-auth" - "https://${NETBIRD_DOMAIN}/nb-silent-auth" cliRedirectURIs: - "http://localhost:53000/" + # Trust X-Forwarded-* only from the Traefik container's static address. Both + # keys must stay in step with the ipv4_address pinned in docker-compose.yml: + # trustedPeers decides whether forwarded headers are read at all, and leaving + # it unset falls back to 0.0.0.0/0. + reverseProxy: + trustedPeers: + - "${TRAEFIK_IP}/32" + trustedHTTPProxies: + - "${TRAEFIK_IP}/32" + store: engine: "postgres" dsn: "${POSTGRES_DSN}" diff --git a/infrastructure_files/getting-started.sh b/infrastructure_files/getting-started.sh index 0206c269a..4f2c1d82e 100755 --- a/infrastructure_files/getting-started.sh +++ b/infrastructure_files/getting-started.sh @@ -348,6 +348,7 @@ initialize_default_values() { NETBIRD_RELAY_AUTH_SECRET=$(openssl rand -base64 32 | sed "$SED_STRIP_PADDING") # Note: DataStoreEncryptionKey must keep base64 padding (=) for Go's base64.StdEncoding DATASTORE_ENCRYPTION_KEY=$(openssl rand -base64 32) + SESSION_COOKIE_ENCRYPTION_KEY=$(openssl rand -base64 32) NETBIRD_STUN_PORT=3478 # Docker images @@ -527,7 +528,8 @@ generate_configuration_files() { # Common files for all configurations render_dashboard_env > dashboard.env - render_combined_yaml > config.yaml + install -m 600 /dev/null config.yaml + render_combined_yaml >> config.yaml return 0 } @@ -911,6 +913,7 @@ server: auth: issuer: "$NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN/oauth2" signKeyRefreshEnabled: true + sessionCookieEncryptionKey: "$SESSION_COOKIE_ENCRYPTION_KEY" dashboardRedirectURIs: - "$NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN/nb-auth" - "$NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN/nb-silent-auth" diff --git a/infrastructure_files/migrate-to-enterprise.sh b/infrastructure_files/migrate-to-enterprise.sh index d4c59699b..8e1fcf521 100755 --- a/infrastructure_files/migrate-to-enterprise.sh +++ b/infrastructure_files/migrate-to-enterprise.sh @@ -15,7 +15,11 @@ set -o pipefail # 2. Postgres migration — add Postgres, migrate SQLite data via migrate-store. # 3. Traffic flow — add NATS + flow-enricher + flow-receiver. # -# To revert: +# If any step fails once the stack has been touched, the script rolls itself +# back automatically: generated files are removed, the Postgres volume this run +# created is dropped, and the original deployment is started again. +# +# To revert a successful migration: # docker compose down # rm -f docker-compose.override.yml config.yaml.enterprise # # If Postgres migration was done, also restore the SQLite backup printed @@ -25,6 +29,15 @@ set -o pipefail OVERRIDE_FILE="docker-compose.override.yml" ENTERPRISE_CONFIG_FILE="config.yaml.enterprise" +# Rollback bookkeeping. ROLLBACK_STATE flips to "armed" the moment the script +# starts mutating the deployment, and back to "disarmed" once the migration has +# completed successfully. +ROLLBACK_STATE="disarmed" +ENV_EXISTED="unknown" +ENV_BACKUP="" +PG_VOLUME_NAME="" +BACKUP_DIR="" + NETBIRD_EULA_URL="https://netbird.io/self-hosted-EULA" check_docker_compose() { @@ -361,7 +374,77 @@ render_enterprise_config() { # Execution steps # --------------------------------------------------------------------------- -resolve_data_volume() { +combined_container_id() { + $DOCKER_COMPOSE_COMMAND ps -aq "$COMBINED_SERVICE" 2>/dev/null | head -1 +} + +container_data_mount() { + local container="$1" + [[ -n "$container" ]] || return 0 + docker inspect "$container" --format \ + '{{range .Mounts}}{{if eq .Destination "/var/lib/netbird"}}{{if .Name}}{{.Name}}{{else}}{{.Source}}{{end}}{{end}}{{end}}' 2>/dev/null +} + +# The name comes from the container, so `-v` cannot invent an empty volume here. +# 0 = empty, 1 = holds data, 2 = could not determine. A failed listing must not +# be reported as empty: that would abort a healthy migration over a pull error +# or an unreadable bind mount. +data_dir_state() { + local src="$1" out + if [[ "$src" == /* ]]; then + [[ -d "$src" ]] || return 2 + out=$(ls -A "$src" 2>/dev/null) || return 2 + else + docker volume inspect "$src" &> /dev/null || return 0 + out=$(docker run --rm -v "${src}:/d:ro" busybox sh -c 'ls -A /d' 2>/dev/null) || return 2 + fi + [[ -z "$out" ]] && return 0 + return 1 +} + +check_data_directory() { + [[ "$MIGRATE_POSTGRES" == "yes" ]] || return 0 + + local container + container=$(combined_container_id) + if [[ -z "$container" ]]; then + echo "" > /dev/stderr + echo "No container found for service '$COMBINED_SERVICE'." > /dev/stderr + echo "The migration backs up the store by copying it out of that container," > /dev/stderr + echo "so it has to exist. Start the deployment and re-run:" > /dev/stderr + echo " $DOCKER_COMPOSE_COMMAND up -d" > /dev/stderr + exit 1 + fi + + local src + src=$(container_data_mount "$container") + if [[ -z "$src" ]]; then + echo "" > /dev/stderr + echo "The '$COMBINED_SERVICE' container has nothing mounted at /var/lib/netbird." > /dev/stderr + echo "Cannot locate the NetBird store to back it up." > /dev/stderr + exit 1 + fi + + local state=0 + data_dir_state "$src" || state=$? + if [[ $state -eq 0 ]]; then + echo "" > /dev/stderr + echo "The NetBird data directory is empty:" > /dev/stderr + echo " $src" > /dev/stderr + echo "There is nothing to migrate. Check that you are running this from the" > /dev/stderr + echo "deployment directory of the NetBird install you mean to migrate." > /dev/stderr + exit 1 + fi + if [[ $state -eq 2 ]]; then + echo " ⚠ Could not read $src to confirm it holds data — continuing." > /dev/stderr + echo " The backup step still fails loudly if it turns out to be empty." > /dev/stderr + fi + + echo " Data directory: $src" +} + +# Only for the Postgres volume, which has no container to read it off yet. +resolve_compose_volume() { local short="$1" local actual # Resolve project-prefixed volume name from Docker Compose config first. @@ -391,18 +474,21 @@ resolve_data_volume() { backup_sqlite() { BACKUP_DIR="$(pwd)/backups/sqlite-pre-enterprise-$(date +%Y%m%d-%H%M%S)" mkdir -p "$BACKUP_DIR" - local data_volume_actual - data_volume_actual=$(resolve_data_volume "$DATA_VOLUME") - echo "Backing up SQLite store from volume '$data_volume_actual' to $BACKUP_DIR ..." - docker run --rm \ - -v "${data_volume_actual}:/var/lib/netbird:ro" \ - -v "${BACKUP_DIR}:/backup" \ - busybox \ - sh -c 'cp -a /var/lib/netbird/. /backup/ 2>/dev/null || true' + + local container + container=$(combined_container_id) + if [[ -z "$container" ]]; then + echo " ⚠ No container found for '$COMBINED_SERVICE' — cannot back up the store." > /dev/stderr + exit 1 + fi + + echo "Backing up the NetBird store to $BACKUP_DIR ..." + docker cp "${container}:/var/lib/netbird/." "$BACKUP_DIR/" + local copied copied=$(find "$BACKUP_DIR" -mindepth 1 | head -1) if [[ -z "$copied" ]]; then - echo " ⚠ Backup directory is empty — the volume '$data_volume_actual' didn't contain data. Aborting." > /dev/stderr + echo " ⚠ Backup directory is empty — /var/lib/netbird held no data. Aborting." > /dev/stderr exit 1 fi echo " done" @@ -414,6 +500,135 @@ run_migrate_store() { echo " done" } +# --------------------------------------------------------------------------- +# Rollback — a failed run must not leave the operator with a stopped stack and +# half-written artifacts. +# --------------------------------------------------------------------------- + +# Resolve the name Compose would give the Postgres volume before the override +# exists, so a leftover volume can be spotted up front. +compose_project_name() { + local container project + container=$($DOCKER_COMPOSE_COMMAND ps -aq 2>/dev/null | head -1) + if [[ -n "$container" ]]; then + project=$(docker inspect "$container" \ + --format '{{index .Config.Labels "com.docker.compose.project"}}' 2>/dev/null) + if [[ -n "$project" ]]; then + echo "$project" + return 0 + fi + fi + project=$($DOCKER_COMPOSE_COMMAND config 2>/dev/null | yq eval '.name // ""' - 2>/dev/null) + if [[ -n "$project" ]] && [[ "$project" != "null" ]]; then + echo "$project" + fi + return 0 +} + +postgres_volume_name() { + local project + project=$(compose_project_name) + if [[ -n "$project" ]]; then + echo "${project}_netbird_postgres" + fi + return 0 +} + +# Postgres skips initdb when its data directory is non-empty, so a volume left +# behind by an interrupted run would keep the old password and old contents, +# and migrate-store would fail against it. +check_stale_postgres_volume() { + [[ "$MIGRATE_POSTGRES" == "yes" ]] || return 0 + + PG_VOLUME_NAME=$(postgres_volume_name) + if [[ -z "$PG_VOLUME_NAME" ]]; then + echo "" + echo " ⚠ Could not determine the Compose project name, so a Postgres volume" + echo " left over from an earlier attempt cannot be checked for. If a" + echo " previous run failed, remove it before continuing:" + echo " docker volume ls | grep netbird_postgres" + return 0 + fi + docker volume inspect "$PG_VOLUME_NAME" &> /dev/null || return 0 + + echo "" + echo " ⚠ A Postgres volume from an earlier attempt already exists:" + echo " $PG_VOLUME_NAME" + echo " Postgres does not re-initialise a non-empty data directory, so the" + echo " migration would run against stale credentials and stale data." + local remove + remove=$(read_yes_no " Remove it and continue?" "y") + if [[ "$remove" != "yes" ]]; then + echo "" > /dev/stderr + echo "Aborted. Remove it manually with: docker volume rm $PG_VOLUME_NAME" > /dev/stderr + exit 1 + fi + docker volume rm "$PG_VOLUME_NAME" > /dev/null + echo " Removed." +} + +# Undo whatever this run changed and start the previous deployment again. +rollback() { + ROLLBACK_STATE="done" + + echo "" + echo "──────────────────────────────────────────────────────────────────────" + echo " Migration failed — restoring the previous deployment" + echo "──────────────────────────────────────────────────────────────────────" + + # Resolve while the override is still present; without it Compose no longer + # knows about the Postgres volume. + local pg_volume="$PG_VOLUME_NAME" + if [[ -z "$pg_volume" ]] && [[ "$MIGRATE_POSTGRES" == "yes" ]]; then + pg_volume=$(postgres_volume_name) + fi + + echo "" + echo "Stopping services ..." + $DOCKER_COMPOSE_COMMAND down || true + + echo "Removing generated files ..." + rm -f "$OVERRIDE_FILE" "$ENTERPRISE_CONFIG_FILE" + + # Restore .env to exactly what it was, or remove it if this run created it. + if [[ "$ENV_EXISTED" == "yes" ]] && [[ -f "$ENV_BACKUP" ]]; then + mv -f "$ENV_BACKUP" .env || echo " ⚠ Could not restore .env from $ENV_BACKUP." > /dev/stderr + elif [[ "$ENV_EXISTED" == "no" ]]; then + rm -f .env || true + fi + + # Only ever the volume this run created — never the NetBird data volume. + if [[ -n "$pg_volume" ]] && [[ "$pg_volume" != "null" ]]; then + echo "Removing Postgres volume $pg_volume ..." + docker volume rm "$pg_volume" &> /dev/null || true + fi + + echo "Starting the previous deployment ..." + if ! $DOCKER_COMPOSE_COMMAND up -d; then + echo "" + echo " ⚠ Could not start the previous deployment automatically." > /dev/stderr + echo " Run: $DOCKER_COMPOSE_COMMAND up -d" > /dev/stderr + fi + + echo "" + echo "Rolled back. Your docker-compose.yml, config.yaml and the NetBird data" + echo "volume were never modified." + if [[ -n "$BACKUP_DIR" ]] && [[ -d "$BACKUP_DIR" ]]; then + echo "The SQLite backup taken during this run is kept at:" + echo " $BACKUP_DIR" + fi + echo "──────────────────────────────────────────────────────────────────────" +} + +on_exit() { + local code=$? + trap - EXIT + if [[ $code -ne 0 ]] && [[ "$ROLLBACK_STATE" == "armed" ]]; then + rollback + fi + exit $code +} + # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- @@ -541,9 +756,15 @@ init_migration() { ENABLE_FLOW="no" echo "Step 3 (traffic flow) skipped — requires Postgres." fi + + check_data_directory + check_stale_postgres_volume } apply_changes() { + # From here on a failure must roll the deployment back. + ROLLBACK_STATE="armed" + echo "" echo "Writing $OVERRIDE_FILE ..." install -m 644 /dev/null "$OVERRIDE_FILE" @@ -564,6 +785,14 @@ apply_changes() { # picks it up automatically. echo "Writing .env additions (mode 600) ..." local ENV_FILE=".env" + # Snapshot the operator's .env so a rollback can restore it byte for byte. + if [[ -f "$ENV_FILE" ]]; then + ENV_EXISTED="yes" + ENV_BACKUP="${ENV_FILE}.pre-enterprise-$(date +%Y%m%d-%H%M%S)" + cp -p "$ENV_FILE" "$ENV_BACKUP" + else + ENV_EXISTED="no" + fi touch "$ENV_FILE" chmod 600 "$ENV_FILE" { @@ -592,11 +821,16 @@ apply_changes() { if [[ "$MIGRATE_POSTGRES" == "yes" ]]; then echo "" - echo "Stopping existing services (volumes preserved) ..." - $DOCKER_COMPOSE_COMMAND down + # Stop, but keep the containers: the backup reads the store out of one. + echo "Stopping services so the store is quiescent ..." + $DOCKER_COMPOSE_COMMAND stop backup_sqlite + echo "" + echo "Removing stopped containers (volumes preserved) ..." + $DOCKER_COMPOSE_COMMAND down + echo "" echo "Starting Postgres ..." $DOCKER_COMPOSE_COMMAND up -d postgres @@ -626,6 +860,9 @@ apply_changes() { echo "" echo "Migration complete." + + # Nothing left to undo. + ROLLBACK_STATE="disarmed" } print_summary() { @@ -643,6 +880,7 @@ print_summary() { echo " $OVERRIDE_FILE" [[ "$MIGRATE_POSTGRES" == "yes" ]] && echo " $ENTERPRISE_CONFIG_FILE" echo " .env (license key + secrets, mode 600)" + [[ "$ENV_EXISTED" == "yes" ]] && [[ -f "$ENV_BACKUP" ]] && echo " $ENV_BACKUP (.env as it was before this run)" [[ "$MIGRATE_POSTGRES" == "yes" ]] && echo " backups/sqlite-pre-enterprise-*/ (SQLite backup)" echo "" echo " Tail logs:" @@ -651,19 +889,27 @@ print_summary() { echo "──────────────────────────────────────────────────────────────────────" echo " To revert" echo "──────────────────────────────────────────────────────────────────────" - echo " $DOCKER_COMPOSE_COMMAND down" if [[ "$MIGRATE_POSTGRES" == "yes" ]]; then - # Resolve project-prefixed volume names now (before override is removed). - local pg_volume data_volume_actual - pg_volume=$(resolve_data_volume "netbird_postgres") - data_volume_actual=$(resolve_data_volume "$DATA_VOLUME") - echo " # Remove the Postgres volume FIRST, before deleting the override file:" - echo " docker volume rm $pg_volume" + # Resolve the project-prefixed volume name now, before the override is gone. + local pg_volume + pg_volume=$(resolve_compose_volume "netbird_postgres") + echo " # Stop, but keep the containers so the store can be copied back in:" + echo " $DOCKER_COMPOSE_COMMAND stop" echo " # Restore SQLite from the backup created during this run:" - echo " docker run --rm -v ${data_volume_actual}:/var/lib/netbird -v ${BACKUP_DIR}:/backup busybox sh -c 'cp -a /backup/. /var/lib/netbird/'" + echo " docker cp ${BACKUP_DIR}/. \$($DOCKER_COMPOSE_COMMAND ps -aq $COMBINED_SERVICE):/var/lib/netbird/" + echo " $DOCKER_COMPOSE_COMMAND down" + echo " docker volume rm $pg_volume" + else + echo " $DOCKER_COMPOSE_COMMAND down" fi echo " rm -f $OVERRIDE_FILE $ENTERPRISE_CONFIG_FILE" - echo " # Remove migrate-to-enterprise.sh additions from .env (search for the timestamp marker)" + if [[ "$ENV_EXISTED" == "yes" ]] && [[ -f "$ENV_BACKUP" ]]; then + echo " mv $ENV_BACKUP .env # restores .env as it was before this run" + elif [[ "$ENV_EXISTED" == "no" ]]; then + echo " rm -f .env # created by this run" + else + echo " # Remove migrate-to-enterprise.sh additions from .env (search for the timestamp marker)" + fi echo " $DOCKER_COMPOSE_COMMAND up -d" echo "──────────────────────────────────────────────────────────────────────" } @@ -672,6 +918,10 @@ print_summary() { # Run # --------------------------------------------------------------------------- +trap on_exit EXIT +# Turn signals into a normal exit so the EXIT trap can roll back. +trap 'exit 130' INT TERM + init_migration apply_changes print_summary diff --git a/management/internals/modules/agentnetwork/handlers/handlers_test.go b/management/internals/modules/agentnetwork/handlers/handlers_test.go index 27ebea5dd..9d855c05d 100644 --- a/management/internals/modules/agentnetwork/handlers/handlers_test.go +++ b/management/internals/modules/agentnetwork/handlers/handlers_test.go @@ -17,6 +17,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/account" nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/store" @@ -61,10 +62,23 @@ func newAgentNetworkHandlerFixture(t *testing.T) *agentNetworkHandlerFixture { Return(true, context.Background(), nil). AnyTimes() - manager := agentnetwork.NewManager(st, perms, nil, nil) + // Swallow activity events so the mutation paths (create/update/delete) + // are exercisable through the HTTP layer. + accounts := account.NewMockManager(ctrl) + accounts.EXPECT(). + StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). + AnyTimes() + accounts.EXPECT(). + UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()). + AnyTimes() + + manager := agentnetwork.NewManager(st, perms, accounts, nil) h := &handler{manager: manager} router := mux.NewRouter() + router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST") + router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET") + router.HandleFunc("/agent-network/providers/{providerId}", h.updateProvider).Methods("PUT") h.addPolicyEndpoints(router) h.addConsumptionEndpoints(router) h.addBudgetRuleEndpoints(router) diff --git a/management/internals/modules/agentnetwork/handlers/providers_handler_test.go b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go index 649224c02..05024cde9 100644 --- a/management/internals/modules/agentnetwork/handlers/providers_handler_test.go +++ b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go @@ -1,7 +1,9 @@ package handlers import ( + "encoding/json" "math" + nethttp "net/http" "testing" "github.com/stretchr/testify/assert" @@ -51,3 +53,50 @@ func TestValidate_ModelRates(t *testing.T) { assert.Error(t, validate(base(m), true), "case %q must be rejected", name) } } + +// TestProviderHandler_UpdateReplacesFullState pins the update contract shared +// with the other PUT endpoints: the request replaces the provider's mutable +// state, so optional fields absent from the JSON land as their zero values. +// The two exceptions are server-side: the api_key (a secret — omitted means +// "not rotated") and the session keypair, both preserved by the manager. The +// identity headers stay on the wire as explicit empty strings so a cleared +// value round-trips. +func TestProviderHandler_UpdateReplacesFullState(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + create := `{ + "provider_id": "openai_api", + "name": "openai", + "upstream_url": "https://api.openai.com", + "api_key": "sk-test", + "enabled": true, + "metadata_disabled": true, + "skip_tls_verification": true, + "extra_values": {"x-portkey-config": "pc-prod-3f2a"}, + "identity_header_user_id": "x-bf-dim-netbird_user_id", + "models": [{"id": "gpt-4o", "input_per_1k": 0.0025, "output_per_1k": 0.01}] + }` + rec := f.do(t, nethttp.MethodPost, "/agent-network/providers", create) + require.Equal(t, nethttp.StatusOK, rec.Code, "create must succeed: %s", rec.Body.String()) + + var created api.AgentNetworkProvider + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &created)) + + // Minimal update: only the required fields, no api_key. Everything + // optional must land as its zero value. + update := `{"provider_id": "openai_api", "name": "openai-renamed", "upstream_url": "https://api.openai.com", "enabled": true}` + rec = f.do(t, nethttp.MethodPut, "/agent-network/providers/"+created.Id, update) + require.Equal(t, nethttp.StatusOK, rec.Code, "update without api_key must succeed (key is preserved): %s", rec.Body.String()) + + var updated api.AgentNetworkProvider + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &updated)) + assert.Equal(t, "openai-renamed", updated.Name, "sent field must apply") + assert.True(t, updated.Enabled, "sent field must apply") + assert.False(t, updated.MetadataDisabled, "omitted metadata_disabled must land as false — PUT replaces the full state") + assert.False(t, updated.SkipTlsVerification, "omitted skip_tls_verification must land as false") + assert.Nil(t, updated.ExtraValues, "omitted extra_values must be cleared") + assert.Equal(t, "", updated.IdentityHeaderUserId, "omitted identity header must be cleared yet stay on the wire") + assert.Empty(t, updated.Models, "omitted models must be cleared") + assert.Contains(t, rec.Body.String(), `"identity_header_user_id":""`, + "cleared identity header must round-trip as an explicit empty string") +} diff --git a/management/internals/modules/agentnetwork/handlers/settings_handler.go b/management/internals/modules/agentnetwork/handlers/settings_handler.go index c65efad0f..171750838 100644 --- a/management/internals/modules/agentnetwork/handlers/settings_handler.go +++ b/management/internals/modules/agentnetwork/handlers/settings_handler.go @@ -2,7 +2,6 @@ package handlers import ( "encoding/json" - "errors" "net/http" "github.com/gorilla/mux" @@ -11,19 +10,20 @@ import ( nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/shared/management/http/api" "github.com/netbirdio/netbird/shared/management/http/util" - "github.com/netbirdio/netbird/shared/management/status" ) // addSettingsEndpoints registers the Agent Network settings routes. The -// settings row is bootstrapped server-side on first provider create; GET reads -// it and PUT updates the mutable collection toggles (cluster/subdomain stay -// immutable). +// settings row is bootstrapped server-side on first provider create or on the +// first PUT carrying a cluster; GET reads it and PUT applies a partial update +// of the mutable collection toggles (cluster/subdomain stay immutable). func (h *handler) addSettingsEndpoints(router *mux.Router) { router.HandleFunc("/agent-network/settings", h.getSettings).Methods("GET", "OPTIONS") router.HandleFunc("/agent-network/settings", h.updateSettings).Methods("PUT", "OPTIONS") } -// updateSettings applies the collection toggles to the account's settings row. +// updateSettings replaces the mutable settings fields on the account's row. +// A request carrying a cluster bootstraps the row when the account doesn't +// have one yet. func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) { userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) if err != nil { @@ -48,11 +48,9 @@ func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) { util.WriteJSONObject(r.Context(), w, updated.ToAPIResponse()) } -// getSettings returns the account's agent-network settings. The settings -// row is bootstrapped on first provider create, so freshly-onboarded -// accounts have nothing to read. Rather than 404-ing in that case (which -// the dashboard would have to special-case), return a JSON null with 200 -// so consumers can branch on the body alone. +// getSettings returns the account's agent-network settings. Accounts that +// haven't been bootstrapped yet read as the defaults with an empty cluster, +// subdomain and endpoint; the manager synthesises that view. func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) { userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) if err != nil { @@ -62,11 +60,6 @@ func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) { settings, err := h.manager.GetSettings(r.Context(), userAuth.AccountId, userAuth.UserId) if err != nil { - var sErr *status.Error - if errors.As(err, &sErr) && sErr.Type() == status.NotFound { - util.WriteJSONObject(r.Context(), w, nil) - return - } util.WriteError(r.Context(), err, w) return } diff --git a/management/internals/modules/agentnetwork/handlers/settings_handler_test.go b/management/internals/modules/agentnetwork/handlers/settings_handler_test.go new file mode 100644 index 000000000..636ec5b26 --- /dev/null +++ b/management/internals/modules/agentnetwork/handlers/settings_handler_test.go @@ -0,0 +1,137 @@ +package handlers + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// TestSettingsHandler_GetUnbootstrappedReturnsDefaults pins the settings-read +// convention shared with the account and DNS settings endpoints: settings +// always read as a JSON object. Before bootstrap that object carries the +// defaults with an empty cluster/subdomain/endpoint (the "not bootstrapped" +// signal) and no timestamps — never a 404 and never the legacy null body. +func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodGet, "/agent-network/settings", "") + require.Equal(t, http.StatusOK, rec.Code, + "unbootstrapped account must read as 200 with defaults: got %d body=%s", rec.Code, rec.Body.String()) + require.NotEqual(t, "null", trimSpace(rec.Body.String()), + "the legacy 200+null shape must not come back") + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.Empty(t, got.Cluster, "cluster must be empty until bootstrapped") + assert.Empty(t, got.Subdomain, "subdomain must be empty until bootstrapped") + assert.Empty(t, got.Endpoint, "endpoint must be empty until bootstrapped, not a bare dot") + assert.True(t, got.EnableLogCollection, "defaults must show log collection on, matching bootstrap") + assert.False(t, got.EnablePromptCollection, "defaults must show prompt collection off") + assert.False(t, got.RedactPii, "defaults must show redaction off") + require.NotNil(t, got.AccessLogRetentionDays) + assert.Equal(t, 30, *got.AccessLogRetentionDays, "defaults must show the bootstrap retention") + assert.Nil(t, got.CreatedAt, "no timestamps before a row exists") + assert.Nil(t, got.UpdatedAt, "no timestamps before a row exists") +} + +// TestSettingsHandler_PutBootstrapsWithCluster covers the settings-first +// bootstrap path: a PUT carrying a cluster on an unbootstrapped account +// creates the row (cluster pinned, subdomain assigned) and applies the +// mutable fields from the same request. +func TestSettingsHandler_PutBootstrapsWithCluster(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": false, "access_log_retention_days": 30}`) + require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String()) + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be pinned from the request") + assert.NotEmpty(t, got.Subdomain, "subdomain must be assigned at bootstrap") + assert.Equal(t, got.Subdomain+".eu.proxy.netbird.io", got.Endpoint, "endpoint must combine subdomain and cluster") + assert.True(t, got.EnableLogCollection, "toggle from the bootstrap request must apply") + assert.True(t, got.EnablePromptCollection, "toggle from the bootstrap request must apply") + require.NotNil(t, got.AccessLogRetentionDays) + assert.Equal(t, 30, *got.AccessLogRetentionDays, "retention from the bootstrap request must apply") + + // The row is now readable via GET. + rec = f.do(t, http.MethodGet, "/agent-network/settings", "") + require.Equal(t, http.StatusOK, rec.Code, "GET after bootstrap must succeed") +} + +// TestSettingsHandler_PutWithoutClusterOnUnbootstrapped pins that a PUT +// without a cluster cannot conjure a settings row out of nothing — there is +// no cluster to pin — and surfaces as 404 like the GET. +func TestSettingsHandler_PutWithoutClusterOnUnbootstrapped(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false}`) + assert.Equal(t, http.StatusNotFound, rec.Code, + "cluster-less PUT on an unbootstrapped account must 404: got %d body=%s", rec.Code, rec.Body.String()) + assert.Contains(t, rec.Body.String(), "cluster", + "the error must point the caller at the bootstrap paths: %s", rec.Body.String()) +} + +// TestSettingsHandler_PutReplacesMutableFields pins the update contract shared +// with the other PUT endpoints: the request replaces every mutable field, so a +// toggle absent from the JSON lands as its zero value rather than being +// preserved. Cluster and subdomain survive untouched. +func TestSettingsHandler_PutReplacesMutableFields(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": true, "access_log_retention_days": 14}`) + require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String()) + + var before api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before)) + + rec = f.do(t, http.MethodPut, "/agent-network/settings", + `{"enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`) + require.Equal(t, http.StatusOK, rec.Code, "update PUT must succeed: %s", rec.Body.String()) + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.True(t, got.EnableLogCollection, "sent toggle must apply") + assert.False(t, got.EnablePromptCollection, "sent toggle must apply") + assert.False(t, got.RedactPii, "sent toggle must apply") + require.NotNil(t, got.AccessLogRetentionDays) + assert.Equal(t, 0, *got.AccessLogRetentionDays, + "retention absent from the request must land as the zero value — PUT replaces all mutable fields") + assert.Equal(t, before.Cluster, got.Cluster, "cluster must survive updates untouched") + assert.Equal(t, before.Subdomain, got.Subdomain, "subdomain must survive updates untouched") +} + +// TestSettingsHandler_PutRejectsClusterChange pins cluster immutability: once +// assigned, a differing cluster is rejected as a validation error instead of +// being silently ignored, so callers never observe a value other than the one +// they sent. Echoing the assigned cluster back stays valid, which lets +// declarative clients send their full desired state idempotently. +func TestSettingsHandler_PutRejectsClusterChange(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`) + require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String()) + + rec = f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "us.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`) + assert.Equal(t, http.StatusUnprocessableEntity, rec.Code, + "cluster change must be rejected as a validation error: got %d body=%s", rec.Code, rec.Body.String()) + + rec = f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": true}`) + require.Equal(t, http.StatusOK, rec.Code, "echoing the assigned cluster must stay valid: %s", rec.Body.String()) + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be unchanged") + assert.True(t, got.RedactPii, "toggle sent alongside the echoed cluster must apply") +} diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index 77c77ce44..ba2c06826 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -157,14 +157,14 @@ func NewManager( } func (m *managerImpl) GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID) @@ -175,9 +175,14 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid // been created yet; otherwise it is ignored (the cluster is pinned on // Settings and every provider in the account routes through it). func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error) { - if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Create); err != nil { return nil, err } + if strings.TrimSpace(bootstrapCluster) != "" { + if err := m.requireSettingsBootstrapPermission(ctx, provider.AccountID, userID); err != nil { + return nil, err + } + } // An empty api_key would silently produce a synthesised service // that 401s on every upstream request. Surface the misconfiguration @@ -202,7 +207,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide } if strings.TrimSpace(bootstrapCluster) != "" { - if _, err := m.bootstrapSettingsIfNeeded(ctx, provider.AccountID, bootstrapCluster); err != nil { + if _, err := m.bootstrapSettingsIfNeeded(ctx, m.store, provider.AccountID, bootstrapCluster); err != nil { // The provider create has already succeeded; logging the // bootstrap miss matches the plan's PoC behaviour. The synth // path treats a missing settings row as a no-op, and the next @@ -218,7 +223,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide } func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error) { - if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Update); err != nil { return nil, err } @@ -257,7 +262,7 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide } func (m *managerImpl) DeleteProvider(ctx context.Context, accountID, userID, providerID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Delete); err != nil { return err } @@ -306,21 +311,21 @@ func pluralize(n int, singular, plural string) string { } func (m *managerImpl) GetAllPolicies(ctx context.Context, accountID, userID string) ([]*types.Policy, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetPolicy(ctx context.Context, accountID, userID, policyID string) (*types.Policy, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkPolicyByID(ctx, store.LockingStrengthNone, accountID, policyID) } func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) { - if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Create); err != nil { return nil, err } @@ -346,7 +351,7 @@ func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *t } func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) { - if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Update); err != nil { return nil, err } @@ -373,7 +378,7 @@ func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *t } func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, policyID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Delete); err != nil { return err } @@ -393,21 +398,21 @@ func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, polic } func (m *managerImpl) GetAllGuardrails(ctx context.Context, accountID, userID string) ([]*types.Guardrail, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetGuardrail(ctx context.Context, accountID, userID, guardrailID string) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkGuardrailByID(ctx, store.LockingStrengthNone, accountID, guardrailID) } func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Create); err != nil { return nil, err } @@ -429,7 +434,7 @@ func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardr } func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Update); err != nil { return nil, err } @@ -452,7 +457,7 @@ func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardr } func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, guardrailID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Delete); err != nil { return err } @@ -473,7 +478,7 @@ func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, gu // GetAllBudgetRules returns every account-level budget rule for the account. func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID string) ([]*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkBudgetRules(ctx, store.LockingStrengthNone, accountID) @@ -481,7 +486,7 @@ func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID s // GetBudgetRule returns a single account-level budget rule. func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, ruleID string) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkBudgetRuleByID(ctx, store.LockingStrengthNone, accountID, ruleID) @@ -491,7 +496,7 @@ func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, rule // enforced at request time (CheckLLMPolicyLimits), not baked into the synth // proxy config, so no reconcile is needed. func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Create); err != nil { return nil, err } @@ -513,7 +518,7 @@ func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule // UpdateBudgetRule updates an existing account-level budget rule. func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Update); err != nil { return nil, err } @@ -536,7 +541,7 @@ func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule // DeleteBudgetRule removes an account-level budget rule. func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, ruleID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Delete); err != nil { return err } @@ -554,40 +559,83 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r return nil } -// UpdateSettings applies the mutable account-level settings — the collection -// toggles — onto the existing row. Cluster and Subdomain are immutable and are -// preserved from the persisted row regardless of the input. Because the -// collection toggles change the synthesised service config (prompt-capture -// gating, access-log emission), a reconcile is triggered so the proxy and peer -// network maps converge on the new state. +// UpdateSettings replaces the mutable account-level settings — the collection +// toggles and retention — on the account's row. When the account has no +// settings row yet, a non-empty settings.Cluster bootstraps one (same path as +// first provider create); without it the update fails with NotFound. On an +// existing row the cluster and subdomain are immutable: a differing +// settings.Cluster is rejected rather than silently ignored so callers never +// observe a value other than what they sent. Because the collection toggles +// change the synthesised service config (prompt-capture gating, access-log +// emission), a reconcile is triggered so the proxy and peer network maps +// converge on the new state. func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) { - if err := m.requirePermission(ctx, settings.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil { return nil, err } - existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID) + requestedCluster := strings.TrimSpace(settings.Cluster) + + // The row lock from LockingStrengthUpdate only holds for the duration of + // the surrounding transaction, so the read, the cluster-immutability + // check, and the save must share one — otherwise concurrent PUTs could + // interleave between them. + var updated *types.Settings + err := m.store.ExecuteInTransaction(ctx, func(tx store.Store) error { + existing, err := tx.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID) + switch { + case err == nil: + if requestedCluster != "" && requestedCluster != existing.Cluster { + return status.Errorf(status.InvalidArgument, "cluster is immutable once assigned (current: %s)", existing.Cluster) + } + case isNotFound(err): + if requestedCluster == "" { + return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; pass cluster to bootstrap them, or create a provider with bootstrap_cluster set") + } + // Bootstrapping pins the cluster and subdomain — a settings + // create on top of the update the caller already passed, matching + // the gate on the provider-create bootstrap path. + if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Create); err != nil { + return err + } + existing, err = m.bootstrapSettingsIfNeeded(ctx, tx, settings.AccountID, requestedCluster) + if err != nil { + return err + } + default: + return fmt.Errorf("get agent network settings: %w", err) + } + + existing.EnableLogCollection = settings.EnableLogCollection + existing.EnablePromptCollection = settings.EnablePromptCollection + existing.RedactPii = settings.RedactPii + existing.AccessLogRetentionDays = settings.AccessLogRetentionDays + existing.UpdatedAt = time.Now().UTC() + + if err := tx.SaveAgentNetworkSettings(ctx, existing); err != nil { + return fmt.Errorf("save agent network settings: %w", err) + } + updated = existing + return nil + }) if err != nil { - return nil, fmt.Errorf("get agent network settings: %w", err) - } - - existing.EnableLogCollection = settings.EnableLogCollection - existing.EnablePromptCollection = settings.EnablePromptCollection - existing.RedactPii = settings.RedactPii - existing.AccessLogRetentionDays = settings.AccessLogRetentionDays - existing.UpdatedAt = time.Now().UTC() - - if err := m.store.SaveAgentNetworkSettings(ctx, existing); err != nil { - return nil, fmt.Errorf("save agent network settings: %w", err) + return nil, err } m.accountManager.StoreEvent(ctx, userID, settings.AccountID, settings.AccountID, activity.AgentNetworkSettingsUpdated, map[string]any{ - "log_collection": existing.EnableLogCollection, - "prompt_collection": existing.EnablePromptCollection, - "redact_pii": existing.RedactPii, + "log_collection": updated.EnableLogCollection, + "prompt_collection": updated.EnablePromptCollection, + "redact_pii": updated.RedactPii, }) m.reconcile(ctx, settings.AccountID) - return existing, nil + return updated, nil +} + +// isNotFound reports whether err is a status.NotFound error. +func isNotFound(err error) bool { + var sErr *status.Error + return errors.As(err, &sErr) && sErr.Type() == status.NotFound } // validateProviderRefs ensures every destination provider id refers to a @@ -611,14 +659,38 @@ func (m *managerImpl) validateProviderRefs(ctx context.Context, accountID string return nil } -// GetSettings returns the agent-network settings row for the account. -// Returns the underlying status.NotFound when no row has been -// bootstrapped yet (i.e. the account has no providers). +// GetSettings returns the agent-network settings row for the account. When no +// row has been bootstrapped yet, the defaults are returned (without +// persisting) with cluster and subdomain empty — settings always read as an +// object, like the account and DNS settings endpoints. func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Read); err != nil { return nil, err } - return m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + switch { + case err == nil: + return settings, nil + case isNotFound(err): + return types.DefaultSettings(accountID), nil + default: + return nil, err + } +} + +// requireSettingsBootstrapPermission gates the one-time settings bootstrap a +// first provider create performs. Pinning the account's cluster and subdomain +// is a settings write, so it needs the settings permission on top of the +// provider one. No-op once the settings row exists. +func (m *managerImpl) requireSettingsBootstrapPermission(ctx context.Context, accountID, userID string) error { + _, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + if err == nil { + return nil + } + if !isNotFound(err) { + return fmt.Errorf("get agent network settings: %w", err) + } + return m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Create) } // bootstrapSettingsIfNeeded creates the per-account agent-network @@ -626,8 +698,9 @@ func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) // hint the dashboard sends (auto-picked from the active cluster list); // the subdomain is picked from the curated wordlist avoiding // collisions on the same cluster. Idempotent: if a row already exists -// it is returned untouched and the hint is ignored. -func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, providerCluster string) (*types.Settings, error) { +// it is returned untouched and the hint is ignored. st is the store to +// operate on — pass the transaction store when calling from within one. +func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, st store.Store, accountID, providerCluster string) (*types.Settings, error) { if accountID == "" { return nil, fmt.Errorf("bootstrap settings: account id is required") } @@ -635,16 +708,15 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, return nil, fmt.Errorf("bootstrap settings: provider cluster is required") } - existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + existing, err := st.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) if err == nil { return existing, nil } - var sErr *status.Error - if !errors.As(err, &sErr) || sErr.Type() != status.NotFound { + if !isNotFound(err) { return nil, fmt.Errorf("get agent network settings: %w", err) } - siblings, err := m.store.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster) + siblings, err := st.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster) if err != nil { return nil, fmt.Errorf("list agent network settings on cluster: %w", err) } @@ -663,18 +735,12 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, m.labelRngMu.Unlock() now := time.Now().UTC() - settings := &types.Settings{ - AccountID: accountID, - Cluster: providerCluster, - Subdomain: subdomain, - // Logs on by default; usage is collected regardless. Retention bounds - // how long full log rows are kept. - EnableLogCollection: true, - AccessLogRetentionDays: types.DefaultAccessLogRetentionDays, - CreatedAt: now, - UpdatedAt: now, - } - if err := m.store.SaveAgentNetworkSettings(ctx, settings); err != nil { + settings := types.DefaultSettings(accountID) + settings.Cluster = providerCluster + settings.Subdomain = subdomain + settings.CreatedAt = now + settings.UpdatedAt = now + if err := st.SaveAgentNetworkSettings(ctx, settings); err != nil { return nil, fmt.Errorf("save agent network settings: %w", err) } return settings, nil @@ -685,7 +751,7 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, // counter view; permission gate is the same Read role that gates // every other agent-network surface. func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID string) ([]*types.Consumption, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil { return nil, err } return m.store.ListAgentNetworkConsumption(ctx, store.LockingStrengthNone, accountID) @@ -694,7 +760,7 @@ func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID str // ListAccessLogs returns a paginated, server-side-filtered page of // agent-network access logs plus the total count matching the filter. func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLog, int64, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil { return nil, 0, err } return m.store.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID, filter) @@ -704,7 +770,7 @@ func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID stri // agent-network access logs grouped by session, plus the total number of // sessions matching the filter. func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil { return nil, 0, err } return m.store.GetAgentNetworkAccessLogSessions(ctx, store.LockingStrengthNone, accountID, filter) @@ -713,7 +779,7 @@ func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, user // GetUsageOverview returns the filtered usage rows aggregated into time buckets // at the requested granularity, oldest-first. func (m *managerImpl) GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil { return nil, err } rows, err := m.store.GetAgentNetworkUsageRows(ctx, store.LockingStrengthNone, accountID, filter) @@ -787,8 +853,8 @@ func (m *managerImpl) RecordConsumption(ctx context.Context, accountID string, k return m.store.IncrementAgentNetworkConsumption(ctx, accountID, kind, dimID, windowSeconds, windowStart, tokensIn, tokensOut, costUSD) } -func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, op operations.Operation) error { - ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetwork, op) +func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, module modules.Module, op operations.Operation) error { + ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, module, op) if err != nil { return status.NewPermissionValidationError(err) } @@ -877,8 +943,8 @@ func (*mockManager) UpdateBudgetRule(_ context.Context, _ string, r *types.Accou func (*mockManager) DeleteBudgetRule(_ context.Context, _, _, _ string) error { return nil } -func (*mockManager) GetSettings(_ context.Context, _, _ string) (*types.Settings, error) { - return nil, status.Errorf(status.NotFound, "agent network settings not found") +func (*mockManager) GetSettings(_ context.Context, accountID, _ string) (*types.Settings, error) { + return types.DefaultSettings(accountID), nil } func (*mockManager) UpdateSettings(_ context.Context, _ string, s *types.Settings) (*types.Settings, error) { diff --git a/management/internals/modules/agentnetwork/provider_bootstrap_test.go b/management/internals/modules/agentnetwork/provider_bootstrap_test.go new file mode 100644 index 000000000..1a2904c51 --- /dev/null +++ b/management/internals/modules/agentnetwork/provider_bootstrap_test.go @@ -0,0 +1,134 @@ +package agentnetwork + +import ( + "context" + "runtime" + "testing" + + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/account" + "github.com/netbirdio/netbird/management/server/permissions" + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/store" + nbtypes "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// bootstrapFixture wires a real sqlite store to a gomock permissions manager +// so tests can grant the provider permission while denying (or never +// expecting) the settings one. +type bootstrapFixture struct { + manager Manager + store store.Store + perms *permissions.MockManager +} + +func newBootstrapFixture(t *testing.T) *bootstrapFixture { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("sqlite store not properly supported on Windows yet") + } + t.Setenv("NETBIRD_STORE_ENGINE", string(nbtypes.SqliteStoreEngine)) + + st, cleanUp, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + require.NoError(t, err, "test store setup must succeed") + t.Cleanup(cleanUp) + + ctrl := gomock.NewController(t) + perms := permissions.NewMockManager(ctrl) + + accounts := account.NewMockManager(ctrl) + accounts.EXPECT().StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + accounts.EXPECT().UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + accounts.EXPECT().BufferUpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + return &bootstrapFixture{ + manager: NewManager(st, perms, accounts, nil), + store: st, + perms: perms, + } +} + +func (f *bootstrapFixture) expectPermission(accountID, userID string, module modules.Module, op operations.Operation, allowed bool) { + f.perms.EXPECT(). + ValidateUserPermissions(gomock.Any(), accountID, userID, module, op). + Return(allowed, context.Background(), nil) +} + +func newBootstrapProvider(accountID string) *types.Provider { + p := types.NewProvider(accountID) + p.Name = "openai" + p.UpstreamURL = "https://api.openai.com" + p.APIKey = "sk-test" + p.Enabled = true + return p +} + +// TestCreateProviderBootstrapRequiresSettingsPermission pins the gate on the +// one-time settings bootstrap: creating the first provider with a +// bootstrap_cluster pins the account's cluster and subdomain, which is a +// settings write and must not ride on the providers permission alone. +func TestCreateProviderBootstrapRequiresSettingsPermission(t *testing.T) { + ctx := context.Background() + + t.Run("denied without settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, false) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.Error(t, err, "bootstrap without settings permission must fail") + var sErr *status.Error + require.ErrorAs(t, err, &sErr) + assert.Equal(t, status.PermissionDenied, sErr.Type(), "denial should surface as permission denied") + + providers, err := f.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "account1") + require.NoError(t, err) + assert.Empty(t, providers, "provider must not be persisted when bootstrap is denied") + _, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1") + assert.Error(t, err, "settings row must not be created when bootstrap is denied") + }) + + t.Run("allowed with settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true) + + created, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.NoError(t, err, "bootstrap with both permissions must succeed") + require.NotNil(t, created) + + settings, err := f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1") + require.NoError(t, err, "bootstrap must create the settings row") + assert.Equal(t, "cluster1.example.com", settings.Cluster, "settings should pin the bootstrap cluster") + }) + + t.Run("existing settings need no settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + require.NoError(t, f.store.SaveAgentNetworkSettings(ctx, &types.Settings{ + AccountID: "account1", + Cluster: "cluster1.example.com", + Subdomain: "existing", + }), "pre-existing settings row setup must succeed") + + // Only the providers permission may be consulted: gomock fails the + // test on any unexpected settings-permission call. + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.NoError(t, err, "create with existing settings must not require the settings permission") + }) + + t.Run("no bootstrap cluster needs no settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "") + require.NoError(t, err, "create without bootstrap must not require the settings permission") + }) +} diff --git a/management/internals/modules/agentnetwork/types/provider.go b/management/internals/modules/agentnetwork/types/provider.go index 96242f45f..b9a194bf6 100644 --- a/management/internals/modules/agentnetwork/types/provider.go +++ b/management/internals/modules/agentnetwork/types/provider.go @@ -164,9 +164,7 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) { p.MetadataDisabled = *req.MetadataDisabled } // Identity-header overrides for catalogs flagged Customizable. - // nil pointer = "field omitted on the wire" → leave the stored - // value untouched (per the openapi description). Empty string is - // an explicit clear that disables stamping for this dimension. + // Empty or omitted disables stamping for this dimension. if req.IdentityHeaderUserId != nil { p.IdentityHeaderUserID = strings.TrimSpace(*req.IdentityHeaderUserId) } @@ -192,16 +190,20 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { created := p.CreatedAt updated := p.UpdatedAt resp := &api.AgentNetworkProvider{ - Id: p.ID, - ProviderId: p.ProviderID, - Name: p.Name, - UpstreamUrl: p.UpstreamURL, - Models: models, - Enabled: p.Enabled, - SkipTlsVerification: p.SkipTLSVerification, - MetadataDisabled: p.MetadataDisabled, - CreatedAt: &created, - UpdatedAt: &updated, + Id: p.ID, + ProviderId: p.ProviderID, + Name: p.Name, + UpstreamUrl: p.UpstreamURL, + Models: models, + // Always present on the wire so an explicitly cleared header + // round-trips as "" instead of vanishing from the response. + IdentityHeaderUserId: p.IdentityHeaderUserID, + IdentityHeaderGroups: p.IdentityHeaderGroups, + Enabled: p.Enabled, + SkipTlsVerification: p.SkipTLSVerification, + MetadataDisabled: p.MetadataDisabled, + CreatedAt: &created, + UpdatedAt: &updated, } if len(p.ExtraValues) > 0 { out := make(map[string]string, len(p.ExtraValues)) @@ -210,14 +212,6 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { } resp.ExtraValues = &out } - if p.IdentityHeaderUserID != "" { - v := p.IdentityHeaderUserID - resp.IdentityHeaderUserId = &v - } - if p.IdentityHeaderGroups != "" { - v := p.IdentityHeaderGroups - resp.IdentityHeaderGroups = &v - } return resp } diff --git a/management/internals/modules/agentnetwork/types/provider_test.go b/management/internals/modules/agentnetwork/types/provider_test.go index f9756bb8b..fd553ece2 100644 --- a/management/internals/modules/agentnetwork/types/provider_test.go +++ b/management/internals/modules/agentnetwork/types/provider_test.go @@ -77,3 +77,41 @@ func TestProvider_MetadataDisabled_RoundTrip(t *testing.T) { assert.False(t, p.MetadataDisabled, "explicit false must clear metadata_disabled") assert.False(t, p.ToAPIResponse().MetadataDisabled, "response must reflect the cleared value") } + +// TestProvider_IdentityHeaders_AlwaysOnWire pins that the identity header +// fields are always present in the API response — an explicitly cleared +// ("") header must round-trip as "" rather than vanish, so API consumers +// (e.g. the Terraform provider) never observe a value other than the one +// they wrote. +func TestProvider_IdentityHeaders_AlwaysOnWire(t *testing.T) { + set := "x-bf-dim-netbird_user_id" + empty := "" + + base := func() *api.AgentNetworkProviderRequest { + return &api.AgentNetworkProviderRequest{ + ProviderId: "custom", + Name: "bifrost", + UpstreamUrl: "https://bifrost.internal", + } + } + + p := NewProvider("acc-1") + resp := p.ToAPIResponse() + assert.Equal(t, "", resp.IdentityHeaderUserId, "unset header must surface as empty string, not be omitted") + assert.Equal(t, "", resp.IdentityHeaderGroups, "unset header must surface as empty string, not be omitted") + + req := base() + req.IdentityHeaderUserId = &set + p.FromAPIRequest(req) + assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "configured header must round-trip") + + // Omitting the field preserves it. + p.FromAPIRequest(base()) + assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "omitted header must preserve the stored value") + + // An explicit "" clears it AND stays visible on the wire. + req = base() + req.IdentityHeaderUserId = &empty + p.FromAPIRequest(req) + assert.Equal(t, "", p.ToAPIResponse().IdentityHeaderUserId, "cleared header must round-trip as empty string") +} diff --git a/management/internals/modules/agentnetwork/types/settings.go b/management/internals/modules/agentnetwork/types/settings.go index d61d9deff..2c53877b5 100644 --- a/management/internals/modules/agentnetwork/types/settings.go +++ b/management/internals/modules/agentnetwork/types/settings.go @@ -1,6 +1,7 @@ package types import ( + "strings" "time" "github.com/netbirdio/netbird/shared/management/http/api" @@ -42,18 +43,34 @@ type Settings struct { // schema cohesive. func (Settings) TableName() string { return "agent_network_settings" } +// DefaultSettings returns the settings an account observes before its row is +// bootstrapped: log collection on with the default retention, everything else +// off, and no cluster/subdomain assigned yet. Bootstrap persists exactly these +// values plus the assigned cluster and subdomain, so the pre-bootstrap read +// and the freshly bootstrapped row agree. +func DefaultSettings(accountID string) *Settings { + return &Settings{ + AccountID: accountID, + EnableLogCollection: true, + AccessLogRetentionDays: DefaultAccessLogRetentionDays, + } +} + // Endpoint returns the bare hostname agents reach this account at: -// `.`. +// `.`. Empty until both halves are assigned at bootstrap. func (s *Settings) Endpoint() string { + if s.Cluster == "" || s.Subdomain == "" { + return "" + } return s.Subdomain + "." + s.Cluster } -// ToAPIResponse renders the settings as the API representation. +// ToAPIResponse renders the settings as the API representation. The +// timestamps are omitted while zero — a default (not yet bootstrapped) view +// has no persisted row to date. func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings { - created := s.CreatedAt - updated := s.UpdatedAt retention := s.AccessLogRetentionDays - return &api.AgentNetworkSettings{ + resp := &api.AgentNetworkSettings{ Cluster: s.Cluster, Subdomain: s.Subdomain, Endpoint: s.Endpoint(), @@ -61,14 +78,27 @@ func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings { EnablePromptCollection: s.EnablePromptCollection, RedactPii: s.RedactPii, AccessLogRetentionDays: &retention, - CreatedAt: &created, - UpdatedAt: &updated, } + if !s.CreatedAt.IsZero() { + created := s.CreatedAt + resp.CreatedAt = &created + } + if !s.UpdatedAt.IsZero() { + updated := s.UpdatedAt + resp.UpdatedAt = &updated + } + return resp } -// FromAPIRequest applies the mutable settings fields from the request. Cluster -// and Subdomain are immutable and intentionally not touched here. +// FromAPIRequest applies the request onto the receiver. The mutable +// collection fields are always replaced with the request values. Cluster +// participates only in bootstrap and the immutability check (see +// Manager.UpdateSettings); Subdomain is server-assigned and never taken +// from a request. func (s *Settings) FromAPIRequest(req *api.AgentNetworkSettingsRequest) { + if req.Cluster != nil { + s.Cluster = strings.TrimSpace(*req.Cluster) + } s.EnableLogCollection = req.EnableLogCollection s.EnablePromptCollection = req.EnablePromptCollection s.RedactPii = req.RedactPii diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index b06407e15..17ffb7d54 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -182,6 +182,10 @@ func (s *BaseServer) GRPCServer() *grpc.Server { grpc.ChainStreamInterceptor(realip.StreamServerInterceptorOpts(realipOpts...), streamInterceptor, proxyStream), } + // Append interceptors contributed by registered gRPC extensions. These + // run after the built-in chain (ChainUnaryInterceptor is additive). + gRPCOpts = appendExtensionInterceptors(gRPCOpts, s.grpcExtensions) + if s.Config.HttpConfig.LetsEncryptDomain != "" { certManager, err := encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain) if err != nil { @@ -213,6 +217,9 @@ func (s *BaseServer) GRPCServer() *grpc.Server { mgmtProto.RegisterProxyServiceServer(gRPCAPIHandler, s.ReverseProxyGRPCServer()) log.Info("ProxyService registered on gRPC server") + // Register services contributed by external modules via the extension seam. + registerExtensions(gRPCAPIHandler, s.grpcExtensions) + return gRPCAPIHandler }) } diff --git a/management/internals/server/grpc_extension.go b/management/internals/server/grpc_extension.go new file mode 100644 index 000000000..3f257c75e --- /dev/null +++ b/management/internals/server/grpc_extension.go @@ -0,0 +1,74 @@ +package server + +import ( + "context" + + "google.golang.org/grpc" +) + +// GRPCExtension bundles an external module's contribution to the management +// gRPC server: the registration of one or more services onto the shared +// grpc.Server, any server-wide interceptors those services require, and an +// optional shutdown hook. It is a generic extension point with no knowledge of +// any specific service. +type GRPCExtension struct { + // Register is invoked with the shared grpc.Server (as a ServiceRegistrar) + // after the built-in services are registered. It may register any number of + // services. May be nil. + Register func(grpc.ServiceRegistrar) + // UnaryInterceptors are appended to the server's unary interceptor chain, + // running after the built-in interceptors. May be empty. + UnaryInterceptors []grpc.UnaryServerInterceptor + // StreamInterceptors are appended to the server's stream interceptor chain, + // running after the built-in interceptors. May be empty. + StreamInterceptors []grpc.StreamServerInterceptor + // Shutdown, if non-nil, is called once during Stop() with the context + // governing server shutdown, which carries a deadline. The hook MUST + // return promptly and MUST abandon its work once that context is + // cancelled or expires: it runs before the rest of Stop()'s cleanup + // (store, event store, embedded IdP) and before Stop() itself checks the + // context's deadline, so a hook that ignores the context will delay all + // of that cleanup and prevent Stop() from returning on time. May be nil. + Shutdown func(ctx context.Context) +} + +// RegisterGRPCExtension registers a gRPC extension. Call before the gRPC server +// is first built (i.e. before Start); registrations after that have no effect. +func (s *BaseServer) RegisterGRPCExtension(ext GRPCExtension) { + s.grpcExtensions = append(s.grpcExtensions, ext) +} + +// appendExtensionInterceptors appends each extension's interceptors to the gRPC +// server options as additional chained interceptors. grpc.ChainUnaryInterceptor +// and grpc.ChainStreamInterceptor are additive, so the returned options run the +// extension interceptors after any interceptors already present in opts. +func appendExtensionInterceptors(opts []grpc.ServerOption, exts []GRPCExtension) []grpc.ServerOption { + for _, ext := range exts { + if len(ext.UnaryInterceptors) > 0 { + opts = append(opts, grpc.ChainUnaryInterceptor(ext.UnaryInterceptors...)) + } + if len(ext.StreamInterceptors) > 0 { + opts = append(opts, grpc.ChainStreamInterceptor(ext.StreamInterceptors...)) + } + } + return opts +} + +// registerExtensions registers each extension's services onto reg. +func registerExtensions(reg grpc.ServiceRegistrar, exts []GRPCExtension) { + for _, ext := range exts { + if ext.Register != nil { + ext.Register(reg) + } + } +} + +// runExtensionShutdownHooks calls each extension's shutdown hook, if set, +// passing ctx through so hooks can honor its deadline/cancellation. +func runExtensionShutdownHooks(ctx context.Context, exts []GRPCExtension) { + for _, ext := range exts { + if ext.Shutdown != nil { + ext.Shutdown(ctx) + } + } +} diff --git a/management/internals/server/grpc_extension_test.go b/management/internals/server/grpc_extension_test.go new file mode 100644 index 000000000..8f444ca72 --- /dev/null +++ b/management/internals/server/grpc_extension_test.go @@ -0,0 +1,160 @@ +package server + +import ( + "context" + "net" + "sync/atomic" + "testing" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/health" + healthgrpc "google.golang.org/grpc/health/grpc_health_v1" + "google.golang.org/grpc/test/bufconn" +) + +// Test that an extension's interceptors and service registration are actually +// wired onto a real in-process gRPC server via the helpers, and that shutdown +// hooks run. This validates the load-bearing assumption that +// grpc.ChainUnaryInterceptor is additive (extension interceptors run in +// addition to any base chain). +func TestGRPCExtensionAppliedToServer(t *testing.T) { + var unaryCalls atomic.Int32 + var streamShutdownCalled atomic.Bool + + ext := GRPCExtension{ + Register: func(reg grpc.ServiceRegistrar) { + healthgrpc.RegisterHealthServer(reg, health.NewServer()) + }, + UnaryInterceptors: []grpc.UnaryServerInterceptor{ + func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + unaryCalls.Add(1) + return handler(ctx, req) + }, + }, + Shutdown: func(ctx context.Context) { streamShutdownCalled.Store(true) }, + } + exts := []GRPCExtension{ext} + + // Base options mimic GRPCServer(): a pre-existing chain the extension appends to. + var baseUnaryCalls atomic.Int32 + opts := []grpc.ServerOption{ + grpc.ChainUnaryInterceptor(func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + baseUnaryCalls.Add(1) + return handler(ctx, req) + }), + } + opts = appendExtensionInterceptors(opts, exts) + + srv := grpc.NewServer(opts...) + registerExtensions(srv, exts) + + lis := bufconn.Listen(1024 * 1024) + go func() { _ = srv.Serve(lis) }() + t.Cleanup(srv.Stop) + + conn, err := grpc.NewClient("passthrough:///bufnet", + grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { return lis.DialContext(ctx) }), + grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = conn.Close() }) + + _, err = healthgrpc.NewHealthClient(conn).Check(context.Background(), &healthgrpc.HealthCheckRequest{}) + if err != nil { + t.Fatalf("health check via extension-registered service failed: %v", err) + } + if baseUnaryCalls.Load() != 1 { + t.Errorf("base interceptor calls = %d, want 1 (base chain must be preserved)", baseUnaryCalls.Load()) + } + if unaryCalls.Load() != 1 { + t.Errorf("extension interceptor calls = %d, want 1", unaryCalls.Load()) + } + + runExtensionShutdownHooks(context.Background(), exts) + if !streamShutdownCalled.Load() { + t.Error("extension shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookReceivesCallerContext asserts that each hook receives +// a non-nil context and that it is the very same context the caller passed +// in, so hooks can rely on values/deadlines placed on it by Stop(). +func TestGRPCExtensionShutdownHookReceivesCallerContext(t *testing.T) { + type sentinelKey struct{} + want := "shutdown-ctx-sentinel" + ctx := context.WithValue(context.Background(), sentinelKey{}, want) + + var called bool + ext := GRPCExtension{ + Shutdown: func(hookCtx context.Context) { + called = true + if hookCtx == nil { + t.Fatal("hook received a nil context") + } + got, _ := hookCtx.Value(sentinelKey{}).(string) + if got != want { + t.Errorf("hook context sentinel = %q, want %q (not the caller's context)", got, want) + } + }, + } + + runExtensionShutdownHooks(ctx, []GRPCExtension{ext}) + if !called { + t.Fatal("shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookObservesCancellation documents, by test, that +// hooks can honor cancellation/deadlines: a hook given an already-cancelled +// context must see ctx.Err() != nil and a closed Done() channel. +func TestGRPCExtensionShutdownHookObservesCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + var called bool + ext := GRPCExtension{ + Shutdown: func(hookCtx context.Context) { + called = true + if hookCtx.Err() == nil { + t.Error("hook context Err() = nil, want non-nil for a cancelled context") + } + select { + case <-hookCtx.Done(): + default: + t.Error("hook context Done() channel is not closed for a cancelled context") + } + }, + } + + runExtensionShutdownHooks(ctx, []GRPCExtension{ext}) + if !called { + t.Fatal("shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookNilSkipped asserts that an extension +// with a nil Shutdown hook is skipped without panicking, and that hooks for +// other extensions still run. +func TestGRPCExtensionShutdownHookNilSkipped(t *testing.T) { + var called atomic.Bool + exts := []GRPCExtension{ + {Shutdown: nil}, + {Shutdown: func(context.Context) { called.Store(true) }}, + } + + runExtensionShutdownHooks(context.Background(), exts) + if !called.Load() { + t.Error("shutdown hook for non-nil extension was not called") + } +} + +func TestRegisterGRPCExtensionAccumulates(t *testing.T) { + s := &BaseServer{} + s.RegisterGRPCExtension(GRPCExtension{}) + s.RegisterGRPCExtension(GRPCExtension{}) + if len(s.grpcExtensions) != 2 { + t.Fatalf("grpcExtensions len = %d, want 2", len(s.grpcExtensions)) + } +} diff --git a/management/internals/server/server.go b/management/internals/server/server.go index 7fd06d947..22a61bada 100644 --- a/management/internals/server/server.go +++ b/management/internals/server/server.go @@ -68,6 +68,11 @@ type BaseServer struct { proxyAuthClose func() + // grpcExtensions holds additional gRPC services, interceptors, and shutdown + // hooks registered by external modules via RegisterGRPCExtension. Populated + // during boot (single-threaded), consumed by GRPCServer() and Stop(). + grpcExtensions []GRPCExtension + listener net.Listener certManager *autocert.Manager update *version.Update @@ -257,6 +262,7 @@ func (s *BaseServer) Stop() error { s.proxyAuthClose() s.proxyAuthClose = nil } + runExtensionShutdownHooks(ctx, s.grpcExtensions) _ = s.Store().Close(ctx) _ = s.EventStore().Close(ctx) if s.update != nil { diff --git a/management/internals/shared/grpc/components_encoder.go b/management/internals/shared/grpc/components_encoder.go index 7e43cf478..e1a5cae48 100644 --- a/management/internals/shared/grpc/components_encoder.go +++ b/management/internals/shared/grpc/components_encoder.go @@ -61,6 +61,8 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel return &proto.NetworkMapEnvelope{ Payload: &proto.NetworkMapEnvelope_Full{ Full: &proto.NetworkMapComponentsFull{ + Serial: networkSerial(c.Network), + Network: toAccountNetwork(c.Network), PeerConfig: in.PeerConfig, // components.Peers always contains the target peer Peers: []*proto.PeerCompact{toPeerCompact(c.Peers[c.PeerID])}, diff --git a/management/internals/shared/grpc/components_encoder_test.go b/management/internals/shared/grpc/components_encoder_test.go index 100ab0948..f7df82f2f 100644 --- a/management/internals/shared/grpc/components_encoder_test.go +++ b/management/internals/shared/grpc/components_encoder_test.go @@ -758,6 +758,9 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) { assert.Equal(t, "netbird.cloud", full.DnsDomain) assert.Len(t, full.Peers, 1) assert.Empty(t, full.Policies) + require.NotNil(t, full.Network, "client runs Calculate() over the envelope and dereferences Network unconditionally; a nil here would crash the receiver") + assert.Equal(t, "net-empty", full.Network.Identifier) + assert.Equal(t, uint64(9), full.Serial) } func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) { @@ -776,6 +779,12 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) { func emptyNetworkMapComponents() *types.NetworkMapComponents { return types.EmptyNetworkMapComponents( &types.NetworkMapComponents{ - PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}}, + PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}, + Network: &types.Network{ + Identifier: "net-empty", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 9, + }, + }, ) } diff --git a/management/server/agentnetwork_budgetrule_realstack_test.go b/management/server/agentnetwork_budgetrule_realstack_test.go index d17f2e26a..790285b4e 100644 --- a/management/server/agentnetwork_budgetrule_realstack_test.go +++ b/management/server/agentnetwork_budgetrule_realstack_test.go @@ -102,11 +102,20 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t require.NotEmpty(t, before.Subdomain, "subdomain pinned at bootstrap") assert.False(t, before.EnablePromptCollection, "prompt collection defaults off") - // Attempt to flip toggles AND smuggle a different cluster/subdomain — the - // immutable fields must be ignored. + // A cluster different from the one pinned at bootstrap must be rejected + // outright — never silently swapped or ignored. + _, err = mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{ + AccountID: accountID, + Cluster: "attacker.cluster", + EnableLogCollection: true, + }) + require.Error(t, err, "UpdateSettings with a mismatched cluster must fail") + + // Flipping the toggles works with the pinned cluster echoed back (and + // with it omitted); the subdomain is never taken from the request. updated, err := mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{ AccountID: accountID, - Cluster: "attacker.cluster", + Cluster: clusterAddr, Subdomain: "evil", EnableLogCollection: true, EnablePromptCollection: true, diff --git a/management/server/group.go b/management/server/group.go index dab891f2a..33870f25e 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -6,6 +6,8 @@ import ( "fmt" "slices" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/rs/xid" log "github.com/sirupsen/logrus" @@ -744,6 +746,14 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty return &GroupLinkError{"network router", linkedRouter.ID} } + if isLinked, linkedService := isGroupLinkedToReverseProxyService(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"reverse proxy service", linkedService.Domain} + } + + if isLinked, linkedPolicy := isGroupLinkedToAgentNetworkPolicy(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"agent network policy", linkedPolicy.Name} + } + return checkGroupLinkedToSettings(ctx, transaction, group) } @@ -875,6 +885,46 @@ func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store, return false, nil } +// isGroupLinkedToReverseProxyService checks if a group is used as an access group +// of a private reverse proxy service or as a bearer-auth distribution group. +func isGroupLinkedToReverseProxyService(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *service.Service) { + services, err := transaction.GetAccountServices(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving reverse proxy services while checking group linkage: %v", err) + return false, nil + } + + for _, svc := range services { + if svc.Private && slices.Contains(svc.AccessGroups, groupID) { + return true, svc + } + if svc.Auth.BearerAuth != nil && svc.Auth.BearerAuth.Enabled && slices.Contains(svc.Auth.BearerAuth.DistributionGroups, groupID) { + return true, svc + } + } + return false, nil +} + +// isGroupLinkedToAgentNetworkPolicy checks if a group is used as a source group by any +// agent network policy in the account. +func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.Policy) { + policies, err := transaction.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving agent network policies while checking group linkage: %v", err) + return false, nil + } + + for _, policy := range policies { + if policy == nil { + continue + } + if slices.Contains(policy.SourceGroups, groupID) { + return true, policy + } + } + return false, nil +} + // areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers. // It fetches each collection once and checks all groupIDs against them in memory. func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) { diff --git a/management/server/group_test.go b/management/server/group_test.go index 22fda2671..deeec61d5 100644 --- a/management/server/group_test.go +++ b/management/server/group_test.go @@ -18,6 +18,8 @@ import ( "golang.org/x/exp/maps" nbdns "github.com/netbirdio/netbird/dns" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/networks" "github.com/netbirdio/netbird/management/server/networks/resources" @@ -125,6 +127,21 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) { "grp-for-integration", "only service users with admin power can delete integration group", }, + { + "agent network policy", + "grp-for-agent-network-policy", + "agent network policy", + }, + { + "reverse proxy private service access group", + "grp-for-rp-private", + "reverse proxy service", + }, + { + "reverse proxy bearer distribution group", + "grp-for-rp-bearer", + "reverse proxy service", + }, } for _, testCase := range testCases { @@ -218,6 +235,17 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) { groupIDs: []string{"grp-for-integration"}, expectedReasons: []string{"only service users with admin power can delete integration group"}, }, + { + name: "agent network policy", + groupIDs: []string{"grp-for-agent-network-policy"}, + expectedReasons: []string{"agent network policy"}, + }, + { + name: "reverse proxy services", + groupIDs: []string{"grp-for-rp-private", "grp-for-rp-bearer"}, + expectedReasons: []string{"reverse proxy service", "reverse proxy service"}, + expectedNotDeleted: []string{"grp-for-rp-private", "grp-for-rp-bearer"}, + }, { name: "successfully delete multiple groups", groupIDs: []string{"group-1", "group-2"}, @@ -285,6 +313,65 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) { } } +func TestDefaultAccountManager_DeleteGroupUnlinkedFromReverseProxyService(t *testing.T) { + am, _, err := createManager(t) + require.NoError(t, err, "Failed to create account manager") + + _, account, err := initTestGroupAccount(am) + require.NoError(t, err, "Failed to init testing account") + + deletableGroups := []*types.Group{ + { + ID: "grp-rp-bearer-disabled", + AccountID: account.Id, + Name: "Group only in a disabled bearer auth", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + }, + { + ID: "grp-rp-nonprivate-access", + AccountID: account.Id, + Name: "Group only in a non-private service's access groups", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + }, + } + for _, group := range deletableGroups { + require.NoError(t, am.CreateGroup(context.Background(), account.Id, groupAdminUserID, group)) + } + + // Disabled bearer auth and stale access groups on a non-private service + // are inert configuration and must not block group deletion. + services := []*rpservice.Service{ + { + ID: "rp-svc-bearer-disabled", + AccountID: account.Id, + Domain: "bearer-disabled.services.example.com", + Auth: rpservice.AuthConfig{ + BearerAuth: &rpservice.BearerAuthConfig{ + Enabled: false, + DistributionGroups: []string{"grp-rp-bearer-disabled"}, + }, + }, + }, + { + ID: "rp-svc-nonprivate-access", + AccountID: account.Id, + Domain: "nonprivate.services.example.com", + Private: false, + AccessGroups: []string{"grp-rp-nonprivate-access"}, + }, + } + for _, svc := range services { + require.NoError(t, am.Store.CreateService(context.Background(), svc)) + } + + for _, group := range deletableGroups { + err = am.DeleteGroup(context.Background(), account.Id, groupAdminUserID, group.ID) + assert.NoError(t, err, "group %s is not referenced by an active reverse proxy gate and should be deletable", group.ID) + } +} + func TestDefaultAccountManager_DeleteGroupLinkedToFlowGroup(t *testing.T) { am, _, err := createManager(t) require.NoError(t, err) @@ -406,6 +493,30 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t Peers: make([]string, 0), } + groupForAgentNetworkPolicy := &types.Group{ + ID: "grp-for-agent-network-policy", + AccountID: "account-id", + Name: "Group for agent network policies", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + + groupForRPPrivate := &types.Group{ + ID: "grp-for-rp-private", + AccountID: "account-id", + Name: "Group for private reverse proxy service", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + + groupForRPBearer := &types.Group{ + ID: "grp-for-rp-bearer", + AccountID: "account-id", + Name: "Group for bearer reverse proxy service", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + routeResource := &route.Route{ ID: "example route", Groups: []string{groupForRoute.ID}, @@ -461,6 +572,66 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForSetupKeys) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForUsers) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForIntegration) + _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkPolicy) + _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPPrivate) + _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPBearer) + + agentNetworkPolicy := &agentNetworkTypes.Policy{ + ID: "example agent network policy", + AccountID: accountID, + Name: "Example agent network policy", + Enabled: true, + SourceGroups: []string{groupForAgentNetworkPolicy.ID}, + } + if err := am.Store.SaveAgentNetworkPolicy(context.Background(), agentNetworkPolicy); err != nil { + return nil, nil, err + } + + // The decoy services are created first so the linkage check has to scan + // past services that do not reference the groups under test. + rpServices := []*rpservice.Service{ + { + ID: "rp-svc-private-decoy", + AccountID: accountID, + Domain: "private-decoy.services.example.com", + Private: true, + AccessGroups: []string{"unrelated-group"}, + }, + { + ID: "rp-svc-bearer-decoy", + AccountID: accountID, + Domain: "bearer-decoy.services.example.com", + Auth: rpservice.AuthConfig{ + BearerAuth: &rpservice.BearerAuthConfig{ + Enabled: true, + DistributionGroups: []string{"unrelated-group"}, + }, + }, + }, + { + ID: "rp-svc-private", + AccountID: accountID, + Domain: "private.services.example.com", + Private: true, + AccessGroups: []string{groupForRPPrivate.ID}, + }, + { + ID: "rp-svc-bearer", + AccountID: accountID, + Domain: "bearer.services.example.com", + Auth: rpservice.AuthConfig{ + BearerAuth: &rpservice.BearerAuthConfig{ + Enabled: true, + DistributionGroups: []string{groupForRPBearer.ID}, + }, + }, + }, + } + for _, svc := range rpServices { + if err := am.Store.CreateService(context.Background(), svc); err != nil { + return nil, nil, err + } + } acc, err := am.Store.GetAccount(context.Background(), account.Id) if err != nil { diff --git a/management/server/permissions/manager.go b/management/server/permissions/manager.go index 995f234d8..6b9977a86 100644 --- a/management/server/permissions/manager.go +++ b/management/server/permissions/manager.go @@ -82,6 +82,9 @@ func (m *managerImpl) ValidateUserPermissions( return m.ValidateRoleModuleAccess(ctx, accountID, role, module, operation), ctxEnriched, nil } +// ValidateRoleModuleAccess resolves an operation against the role's explicit +// grant for the module, then the grant for its parent module when the module +// is a dotted submodule, and finally the role's AutoAllowNew default. func (m *managerImpl) ValidateRoleModuleAccess( ctx context.Context, accountID string, @@ -89,7 +92,7 @@ func (m *managerImpl) ValidateRoleModuleAccess( module modules.Module, operation operations.Operation, ) bool { - if permissions, ok := role.Permissions[module]; ok { + if permissions, ok := lookupModulePermissions(role, module); ok { if allowed, exists := permissions[operation]; exists { return allowed } @@ -100,6 +103,21 @@ func (m *managerImpl) ValidateRoleModuleAccess( return role.AutoAllowNew[operation] } +// lookupModulePermissions returns the role's explicit permission set for the +// module, falling back to the parent module's set for dotted submodules. The +// second return reports whether any explicit set was found. +func lookupModulePermissions(role roles.RolePermissions, module modules.Module) (map[operations.Operation]bool, bool) { + if permissions, ok := role.Permissions[module]; ok { + return permissions, true + } + if parent, hasParent := module.Parent(); hasParent { + if permissions, ok := role.Permissions[parent]; ok { + return permissions, true + } + } + return nil, false +} + func (m *managerImpl) ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) (context.Context, error) { if user.AccountID != accountID { return ctx, status.NewUserNotPartOfAccountError() @@ -119,7 +137,7 @@ func (m *managerImpl) GetPermissionsByRole(ctx context.Context, role types.UserR permissions := roles.Permissions{} for k := range modules.All { - if rolePermissions, ok := roleMap.Permissions[k]; ok { + if rolePermissions, ok := lookupModulePermissions(roleMap, k); ok { permissions[k] = rolePermissions continue } diff --git a/management/server/permissions/manager_test.go b/management/server/permissions/manager_test.go new file mode 100644 index 000000000..345212f43 --- /dev/null +++ b/management/server/permissions/manager_test.go @@ -0,0 +1,139 @@ +package permissions + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/permissions/roles" + "github.com/netbirdio/netbird/management/server/types" +) + +func TestValidateRoleModuleAccessSubmoduleCascade(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + fullAccess := map[operations.Operation]bool{ + operations.Read: true, + operations.Create: true, + operations.Update: true, + operations.Delete: true, + } + readOnly := map[operations.Operation]bool{ + operations.Read: true, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + } + denyAll := map[operations.Operation]bool{ + operations.Read: false, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + } + + t.Run("parent grant covers submodules", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{modules.AgentNetwork: fullAccess}, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Create), + "parent full grant should allow create on a submodule") + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read), + "parent full grant should allow read on a submodule") + }) + + t.Run("submodule grant does not leak to parent or siblings", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{modules.AgentNetworkUsage: readOnly}, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read), + "explicit submodule read should be allowed") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Create), + "read-only submodule grant should not allow create") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetwork, operations.Read), + "submodule grant should not grant the parent module") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read), + "submodule grant should not grant a sibling submodule") + }) + + t.Run("explicit submodule entry wins over parent grant", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{ + modules.AgentNetwork: fullAccess, + modules.AgentNetworkLogs: denyAll, + }, + } + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read), + "explicit submodule deny should override the parent grant") + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read), + "sibling submodules should still resolve through the parent grant") + }) + + t.Run("auto allow applies when neither submodule nor parent is granted", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: readOnly, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read), + "auto-allow read should apply to submodules") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Delete), + "auto-allow should not grant unlisted operations") + }) +} + +// TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules pins the behavior the +// submodule split must not change: every built-in role resolves the new +// submodules exactly as it resolved the agent_network module before. +func TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + submodules := []modules.Module{ + modules.AgentNetworkProviders, + modules.AgentNetworkPolicies, + modules.AgentNetworkGuardrails, + modules.AgentNetworkBudgets, + modules.AgentNetworkUsage, + modules.AgentNetworkLogs, + modules.AgentNetworkSettings, + } + allOperations := []operations.Operation{operations.Read, operations.Create, operations.Update, operations.Delete} + + for _, role := range []types.UserRole{types.UserRoleOwner, types.UserRoleAdmin, types.UserRoleAuditor, types.UserRoleNetworkAdmin, types.UserRoleUser} { + rolePermissions, ok := roles.RolesMap[role] + require.True(t, ok, "role %s must exist in RolesMap", role) + + for _, sub := range submodules { + for _, op := range allOperations { + expected := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, modules.AgentNetwork, op) + actual := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, sub, op) + assert.Equal(t, expected, actual, "role %s: %s on %s should match the agent_network module", role, op, sub) + } + } + } +} + +func TestGetPermissionsByRoleIncludesSubmodules(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + permissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAuditor) + require.NoError(t, err, "auditor role must resolve") + + usage, ok := permissions[modules.AgentNetworkUsage] + require.True(t, ok, "permissions map should contain the usage submodule") + assert.True(t, usage[operations.Read], "auditor should read the usage submodule") + assert.False(t, usage[operations.Update], "auditor should not update the usage submodule") + + adminPermissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAdmin) + require.NoError(t, err, "admin role must resolve") + providers, ok := adminPermissions[modules.AgentNetworkProviders] + require.True(t, ok, "permissions map should contain the providers submodule") + assert.True(t, providers[operations.Delete], "admin should delete on the providers submodule") +} diff --git a/management/server/permissions/modules/module.go b/management/server/permissions/modules/module.go index a3a9c554d..8a2a1a52d 100644 --- a/management/server/permissions/modules/module.go +++ b/management/server/permissions/modules/module.go @@ -1,5 +1,7 @@ package modules +import "strings" + type Module string const ( @@ -20,6 +22,17 @@ const ( IdentityProviders Module = "identity_providers" Services Module = "services" AgentNetwork Module = "agent_network" + + // Agent Network submodules. A role may grant one of these directly + // or grant the AgentNetwork parent, which covers all of them (see + // permissions.Manager cascade resolution). + AgentNetworkProviders Module = "agent_network.providers" + AgentNetworkPolicies Module = "agent_network.policies" + AgentNetworkGuardrails Module = "agent_network.guardrails" + AgentNetworkBudgets Module = "agent_network.budgets" + AgentNetworkUsage Module = "agent_network.usage" + AgentNetworkLogs Module = "agent_network.logs" + AgentNetworkSettings Module = "agent_network.settings" ) var All = map[Module]struct{}{ @@ -40,4 +53,21 @@ var All = map[Module]struct{}{ IdentityProviders: {}, Services: {}, AgentNetwork: {}, + + AgentNetworkProviders: {}, + AgentNetworkPolicies: {}, + AgentNetworkGuardrails: {}, + AgentNetworkBudgets: {}, + AgentNetworkUsage: {}, + AgentNetworkLogs: {}, + AgentNetworkSettings: {}, +} + +// Parent returns the module owning a dotted submodule name and true, or the +// module itself and false when it has no parent. +func (m Module) Parent() (Module, bool) { + if i := strings.IndexByte(string(m), '.'); i > 0 { + return Module(string(m)[:i]), true + } + return m, false } diff --git a/management/server/types/account.go b/management/server/types/account.go index 1a3a30544..474825281 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -1707,14 +1707,34 @@ func (a *Account) injectPrivateServicePolicies(svc *service.Service, proxyPeers if len(proxyPeers) == 0 { return } + // A service's AccessGroups can name groups that no longer exist — persisted + // services and the agent-network synthesiser both carry the ids verbatim from + // their own state. An unresolvable source authorises nothing, so drop it here + // rather than let the network-map assembly resolve it to a nil group. + sources := a.existingGroupIDs(svc.AccessGroups) + if len(sources) == 0 { + return + } for _, proxyPeer := range proxyPeers { - a.Policies = append(a.Policies, a.createPrivateServicePolicy(svc, proxyPeer)) + a.Policies = append(a.Policies, a.createPrivateServicePolicy(svc, proxyPeer, sources)) } } -func (a *Account) createPrivateServicePolicy(svc *service.Service, proxyPeer *nbpeer.Peer) *Policy { +// existingGroupIDs returns the subset of groupIDs that resolve to a group in the account, +// preserving the input order. +func (a *Account) existingGroupIDs(groupIDs []string) []string { + out := make([]string, 0, len(groupIDs)) + for _, groupID := range groupIDs { + if _, ok := a.Groups[groupID]; ok { + out = append(out, groupID) + } + } + return out +} + +func (a *Account) createPrivateServicePolicy(svc *service.Service, proxyPeer *nbpeer.Peer, accessGroups []string) *Policy { policyID := fmt.Sprintf("private-access-%s-%s", svc.ID, proxyPeer.ID) - sources := append([]string(nil), svc.AccessGroups...) + sources := append([]string(nil), accessGroups...) return &Policy{ ID: policyID, Name: fmt.Sprintf("Private Access to %s", svc.Name), diff --git a/management/server/types/proxy_access_token.go b/management/server/types/proxy_access_token.go index b20b83bc1..9bb27ef02 100644 --- a/management/server/types/proxy_access_token.go +++ b/management/server/types/proxy_access_token.go @@ -68,7 +68,7 @@ type ProxyAccessTokenGenerated struct { // CreateNewProxyAccessToken generates a new proxy access token. // Returns the token with hashed value stored and plain token for one-time display. func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID *string, createdBy string) (*ProxyAccessTokenGenerated, error) { - hashedToken, plainToken, err := generateProxyToken() + hashedToken, plainToken, err := GenerateProxyToken() if err != nil { return nil, err } @@ -94,7 +94,10 @@ func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID * }, nil } -func generateProxyToken() (HashedProxyToken, PlainProxyToken, error) { +// GenerateProxyToken generates a new random proxy token, returning its SHA-256 +// hash (for storage) and the one-time plaintext. Exported so external modules +// can mint tokens in the canonical proxy-token format. +func GenerateProxyToken() (HashedProxyToken, PlainProxyToken, error) { secret, err := b.Random(ProxyTokenSecretLength) if err != nil { return "", "", err diff --git a/management/server/types/proxy_access_token_test.go b/management/server/types/proxy_access_token_test.go index aa1a4d2dd..740b87c2f 100644 --- a/management/server/types/proxy_access_token_test.go +++ b/management/server/types/proxy_access_token_test.go @@ -1,6 +1,7 @@ package types import ( + "strings" "testing" "time" @@ -123,6 +124,22 @@ func TestCreateNewProxyAccessToken(t *testing.T) { }) } +func TestGenerateProxyToken(t *testing.T) { + hashed, plain, err := GenerateProxyToken() + if err != nil { + t.Fatal(err) + } + if err := plain.Validate(); err != nil { + t.Errorf("generated token failed Validate(): %v", err) + } + if plain.Hash() != hashed { + t.Error("returned hashed token does not match Hash(plain)") + } + if !strings.HasPrefix(string(plain), ProxyTokenPrefix) { + t.Errorf("token %q missing prefix %q", plain, ProxyTokenPrefix) + } +} + func TestProxyAccessToken_IsExpired(t *testing.T) { past := time.Now().Add(-1 * time.Hour) future := time.Now().Add(1 * time.Hour) diff --git a/release_files/darwin_pkg/postinstall b/release_files/darwin_pkg/postinstall index 33fa4bfee..2c96a80cd 100755 --- a/release_files/darwin_pkg/postinstall +++ b/release_files/darwin_pkg/postinstall @@ -30,7 +30,23 @@ mkdir -p /usr/local/bin/ $AGENT service install || true $AGENT service start || true - open $APP + console_user=$(stat -f%Su /dev/console 2>/dev/null) + case "$console_user" in + ""|root|loginwindow|_mbsetupuser) + echo "No active GUI user session (console user: '${console_user:-none}'); skipping UI launch." + ;; + *) + uid=$(id -u "$console_user" 2>/dev/null) + if [ -z "$uid" ]; then + echo "Could not resolve uid for console user '$console_user'; skipping UI launch." + else + echo "Launching NetBird UI as console user $console_user (uid $uid)." + if ! launchctl asuser "$uid" sudo -u "$console_user" -H open "$APP"; then + echo "Failed to launch NetBird UI; if autostart is enabled it will start at next login." + fi + fi + ;; + esac echo "Finished Netbird installation successfully" exit 0 # all good diff --git a/shared/management/client/client.go b/shared/management/client/client.go index 8205e3a4f..c48e1ed3e 100644 --- a/shared/management/client/client.go +++ b/shared/management/client/client.go @@ -22,7 +22,6 @@ type Client interface { ExtendAuthSession(sysInfo *system.Info, jwtToken string) (*proto.ExtendAuthSessionResponse, error) GetDeviceAuthorizationFlow() (*proto.DeviceAuthorizationFlow, error) GetPKCEAuthorizationFlow() (*proto.PKCEAuthorizationFlow, error) - GetNetworkMap(sysInfo *system.Info) (*proto.NetworkMap, error) GetServerURL() string // IsHealthy returns the current connection status without blocking. // Used by the engine to monitor connectivity in the background. diff --git a/shared/management/client/grpc.go b/shared/management/client/grpc.go index bd2d0da1f..81f25900a 100644 --- a/shared/management/client/grpc.go +++ b/shared/management/client/grpc.go @@ -436,49 +436,6 @@ func (c *GrpcClient) handleSyncStream(ctx context.Context, serverPubKey wgtypes. return nil } -// GetNetworkMap return with the network map -func (c *GrpcClient) GetNetworkMap(sysInfo *system.Info) (*proto.NetworkMap, error) { - serverPubKey, err := c.getServerPublicKey() - if err != nil { - log.Debugf("failed getting Management Service public key: %s", err) - return nil, err - } - - ctx, cancelStream := context.WithCancel(c.ctx) - defer cancelStream() - stream, err := c.connectToSyncStream(ctx, *serverPubKey, sysInfo) - if err != nil { - log.Debugf("failed to open Management Service stream: %s", err) - return nil, err - } - defer func() { - _ = stream.CloseSend() - }() - - update, err := stream.Recv() - if err == io.EOF { - log.Debugf("Management stream has been closed by server: %s", err) - return nil, err - } - if err != nil { - log.Debugf("disconnected from Management Service sync stream: %v", err) - return nil, err - } - - decryptedResp := &proto.SyncResponse{} - err = encryption.DecryptMessage(*serverPubKey, c.key, update.Body, decryptedResp) - if err != nil { - log.Errorf("failed decrypting update message from Management Service: %s", err) - return nil, err - } - - if decryptedResp.GetNetworkMap() == nil { - return nil, fmt.Errorf("invalid msg, required network map") - } - - return decryptedResp.GetNetworkMap(), nil -} - func (c *GrpcClient) connectToSyncStream(ctx context.Context, serverPubKey wgtypes.Key, sysInfo *system.Info) (proto.ManagementService_SyncClient, error) { req := &proto.SyncRequest{Meta: infoToMetaData(sysInfo)} diff --git a/shared/management/client/mock.go b/shared/management/client/mock.go index ba156a225..e57e314da 100644 --- a/shared/management/client/mock.go +++ b/shared/management/client/mock.go @@ -94,11 +94,6 @@ func (m *MockClient) HealthCheck() error { return m.HealthCheckFunc() } -// GetNetworkMap mock implementation of GetNetworkMap from Client interface. -func (m *MockClient) GetNetworkMap(_ *system.Info) (*proto.NetworkMap, error) { - return nil, nil -} - // GetServerURL mock implementation of GetServerURL from mgm.Client interface func (m *MockClient) GetServerURL() string { if m.GetServerURLFunc == nil { diff --git a/shared/management/client/rest/agentnetwork.go b/shared/management/client/rest/agentnetwork.go new file mode 100644 index 000000000..cee053d17 --- /dev/null +++ b/shared/management/client/rest/agentnetwork.go @@ -0,0 +1,381 @@ +package rest + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// AgentNetworkAPI APIs for the Agent Network (AI/LLM gateway), do not use directly +// see more: https://docs.netbird.io/api/resources/agent-network +type AgentNetworkAPI struct { + c *Client +} + +// ListCatalogProviders lists the catalog of supported upstream AI providers +// (openai_api, anthropic_api, bedrock_api, ...) with their default models and +// pricing, used to prefill provider create forms. +func (a *AgentNetworkAPI) ListCatalogProviders(ctx context.Context) ([]api.AgentNetworkCatalogProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/catalog/providers", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkCatalogProvider](resp) + return ret, err +} + +// ListProviders lists all Agent Network providers +func (a *AgentNetworkAPI) ListProviders(ctx context.Context) ([]api.AgentNetworkProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/providers", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkProvider](resp) + return ret, err +} + +// GetProvider gets Agent Network provider info +func (a *AgentNetworkAPI) GetProvider(ctx context.Context, providerID string) (*api.AgentNetworkProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/providers/"+providerID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// CreateProvider creates a new Agent Network provider. Set +// request.BootstrapCluster on the account's first provider to bootstrap the +// per-account gateway endpoint (alternatively bootstrap via UpdateSettings +// with a cluster). +func (a *AgentNetworkAPI) CreateProvider(ctx context.Context, request api.PostApiAgentNetworkProvidersJSONRequestBody) (*api.AgentNetworkProvider, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/providers", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// UpdateProvider updates an Agent Network provider. The request replaces the +// provider's mutable state; only an omitted api_key keeps the stored key +// (secrets are never required to round-trip). +func (a *AgentNetworkAPI) UpdateProvider(ctx context.Context, providerID string, request api.PutApiAgentNetworkProvidersProviderIdJSONRequestBody) (*api.AgentNetworkProvider, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/providers/"+providerID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// DeleteProvider deletes an Agent Network provider. Fails while any policy +// still references the provider — detach it first. +func (a *AgentNetworkAPI) DeleteProvider(ctx context.Context, providerID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/providers/"+providerID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListPolicies lists all Agent Network policies +func (a *AgentNetworkAPI) ListPolicies(ctx context.Context) ([]api.AgentNetworkPolicy, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/policies", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkPolicy](resp) + return ret, err +} + +// GetPolicy gets Agent Network policy info +func (a *AgentNetworkAPI) GetPolicy(ctx context.Context, policyID string) (*api.AgentNetworkPolicy, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/policies/"+policyID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// CreatePolicy creates a new Agent Network policy +func (a *AgentNetworkAPI) CreatePolicy(ctx context.Context, request api.PostApiAgentNetworkPoliciesJSONRequestBody) (*api.AgentNetworkPolicy, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/policies", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// UpdatePolicy updates an Agent Network policy +func (a *AgentNetworkAPI) UpdatePolicy(ctx context.Context, policyID string, request api.PutApiAgentNetworkPoliciesPolicyIdJSONRequestBody) (*api.AgentNetworkPolicy, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/policies/"+policyID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// DeletePolicy deletes an Agent Network policy +func (a *AgentNetworkAPI) DeletePolicy(ctx context.Context, policyID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/policies/"+policyID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListGuardrails lists all Agent Network guardrails +func (a *AgentNetworkAPI) ListGuardrails(ctx context.Context) ([]api.AgentNetworkGuardrail, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/guardrails", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkGuardrail](resp) + return ret, err +} + +// GetGuardrail gets Agent Network guardrail info +func (a *AgentNetworkAPI) GetGuardrail(ctx context.Context, guardrailID string) (*api.AgentNetworkGuardrail, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/guardrails/"+guardrailID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// CreateGuardrail creates a new Agent Network guardrail +func (a *AgentNetworkAPI) CreateGuardrail(ctx context.Context, request api.PostApiAgentNetworkGuardrailsJSONRequestBody) (*api.AgentNetworkGuardrail, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/guardrails", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// UpdateGuardrail updates an Agent Network guardrail +func (a *AgentNetworkAPI) UpdateGuardrail(ctx context.Context, guardrailID string, request api.PutApiAgentNetworkGuardrailsGuardrailIdJSONRequestBody) (*api.AgentNetworkGuardrail, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/guardrails/"+guardrailID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// DeleteGuardrail deletes an Agent Network guardrail +func (a *AgentNetworkAPI) DeleteGuardrail(ctx context.Context, guardrailID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/guardrails/"+guardrailID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListBudgetRules lists all account-level Agent Network budget rules +func (a *AgentNetworkAPI) ListBudgetRules(ctx context.Context) ([]api.AgentNetworkBudgetRule, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/budget-rules", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkBudgetRule](resp) + return ret, err +} + +// GetBudgetRule gets Agent Network budget rule info +func (a *AgentNetworkAPI) GetBudgetRule(ctx context.Context, ruleID string) (*api.AgentNetworkBudgetRule, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/budget-rules/"+ruleID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// CreateBudgetRule creates a new Agent Network budget rule +func (a *AgentNetworkAPI) CreateBudgetRule(ctx context.Context, request api.PostApiAgentNetworkBudgetRulesJSONRequestBody) (*api.AgentNetworkBudgetRule, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/budget-rules", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// UpdateBudgetRule updates an Agent Network budget rule +func (a *AgentNetworkAPI) UpdateBudgetRule(ctx context.Context, ruleID string, request api.PutApiAgentNetworkBudgetRulesRuleIdJSONRequestBody) (*api.AgentNetworkBudgetRule, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/budget-rules/"+ruleID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// DeleteBudgetRule deletes an Agent Network budget rule +func (a *AgentNetworkAPI) DeleteBudgetRule(ctx context.Context, ruleID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/budget-rules/"+ruleID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// GetSettings gets the account's Agent Network gateway settings (cluster, +// subdomain, endpoint, collection toggles). An account that has not been +// bootstrapped yet — via UpdateSettings with a cluster, or by creating the +// first provider with bootstrap_cluster set — reads as the defaults with an +// empty Cluster, Subdomain and Endpoint. Management servers prior to that +// contract answered 200 with a JSON null body instead; that legacy shape is +// translated to an APIError matchable via IsNotFound rather than fabricating +// defaults the server never stated. +func (a *AgentNetworkAPI) GetSettings(ctx context.Context) (*api.AgentNetworkSettings, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/settings", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, err + } + if trimmed := bytes.TrimSpace(body); len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil, &APIError{StatusCode: http.StatusNotFound, Message: "agent network settings not found"} + } + var ret api.AgentNetworkSettings + if err := json.Unmarshal(body, &ret); err != nil { + return nil, err + } + return &ret, nil +} + +// UpdateSettings updates the account's Agent Network settings; the request +// replaces every mutable field (collection toggles and retention). Setting +// request.Cluster bootstraps the settings row when the account does not have +// one yet; on a bootstrapped account it must match the assigned cluster (or +// be nil) and any other value is rejected — the cluster is immutable. +func (a *AgentNetworkAPI) UpdateSettings(ctx context.Context, request api.PutApiAgentNetworkSettingsJSONRequestBody) (*api.AgentNetworkSettings, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/settings", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkSettings](resp) + return &ret, err +} diff --git a/shared/management/client/rest/agentnetwork_test.go b/shared/management/client/rest/agentnetwork_test.go new file mode 100644 index 000000000..053859125 --- /dev/null +++ b/shared/management/client/rest/agentnetwork_test.go @@ -0,0 +1,497 @@ +//go:build integration + +package rest_test + +import ( + "context" + "encoding/json" + "io" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/client/rest" + "github.com/netbirdio/netbird/shared/management/http/api" + "github.com/netbirdio/netbird/shared/management/http/util" +) + +var ( + testAgentNetworkProvider = api.AgentNetworkProvider{ + Id: "ainp_test", + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + Models: []api.AgentNetworkProviderModel{}, + Enabled: true, + } + + testAgentNetworkPolicy = api.AgentNetworkPolicy{ + Id: "ainpol_test", + Name: "Engineering → OpenAI", + Enabled: true, + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + } + + testAgentNetworkGuardrail = api.AgentNetworkGuardrail{ + Id: "aingr_test", + Name: "No secrets", + } + + testAgentNetworkBudgetRule = api.AgentNetworkBudgetRule{ + Id: "ainbud_test", + Name: "Org monthly ceiling", + Enabled: true, + } + + testAgentNetworkSettings = api.AgentNetworkSettings{ + Cluster: "eu.proxy.netbird.io", + Subdomain: "violet", + Endpoint: "violet.eu.proxy.netbird.io", + EnableLogCollection: true, + AccessLogRetentionDays: ptr(30), + } +) + +func TestAgentNetwork_ListCatalogProviders_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/catalog/providers", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkCatalogProvider{{Id: "openai_api", Name: "OpenAI"}}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListCatalogProviders(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, "openai_api", ret[0].Id) + }) +} + +func TestAgentNetwork_ListProviders_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkProvider{testAgentNetworkProvider}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListProviders(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkProvider, ret[0]) + }) +} + +func TestAgentNetwork_GetProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "GET", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetProvider(context.Background(), "ainp_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_GetProvider_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "not found", Code: 404}) + w.WriteHeader(404) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.GetProvider(context.Background(), "ainp_test") + require.Error(t, err) + assert.True(t, rest.IsNotFound(err), "a 404 must be matchable via IsNotFound") + }) +} + +func TestAgentNetwork_CreateProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + var req api.PostApiAgentNetworkProvidersJSONRequestBody + require.NoError(t, json.Unmarshal(reqBytes, &req)) + assert.Equal(t, "OpenAI", req.Name) + require.NotNil(t, req.BootstrapCluster) + assert.Equal(t, "eu.proxy.netbird.io", *req.BootstrapCluster) + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateProvider(context.Background(), api.PostApiAgentNetworkProvidersJSONRequestBody{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + ApiKey: ptr("sk-test"), + BootstrapCluster: ptr("eu.proxy.netbird.io"), + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_UpdateProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + // Omitted optional fields must be absent from the wire (not + // zero-valued) so the server-side merge preserves them. + assert.NotContains(t, string(reqBytes), "api_key") + assert.NotContains(t, string(reqBytes), "models") + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateProvider(context.Background(), "ainp_test", api.PutApiAgentNetworkProvidersProviderIdJSONRequestBody{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_DeleteProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteProvider(context.Background(), "ainp_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListPolicies_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkPolicy{testAgentNetworkPolicy}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListPolicies(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkPolicy, ret[0]) + }) +} + +func TestAgentNetwork_GetPolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetPolicy(context.Background(), "ainpol_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_CreatePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreatePolicy(context.Background(), api.PostApiAgentNetworkPoliciesJSONRequestBody{ + Name: "Engineering → OpenAI", + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_UpdatePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdatePolicy(context.Background(), "ainpol_test", api.PutApiAgentNetworkPoliciesPolicyIdJSONRequestBody{ + Name: "Engineering → OpenAI", + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_DeletePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeletePolicy(context.Background(), "ainpol_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListGuardrails_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkGuardrail{testAgentNetworkGuardrail}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListGuardrails(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkGuardrail, ret[0]) + }) +} + +func TestAgentNetwork_GetGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetGuardrail(context.Background(), "aingr_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_CreateGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateGuardrail(context.Background(), api.PostApiAgentNetworkGuardrailsJSONRequestBody{ + Name: "No secrets", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_UpdateGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateGuardrail(context.Background(), "aingr_test", api.PutApiAgentNetworkGuardrailsGuardrailIdJSONRequestBody{ + Name: "No secrets", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_DeleteGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteGuardrail(context.Background(), "aingr_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListBudgetRules_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkBudgetRule{testAgentNetworkBudgetRule}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListBudgetRules(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkBudgetRule, ret[0]) + }) +} + +func TestAgentNetwork_GetBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetBudgetRule(context.Background(), "ainbud_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_CreateBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateBudgetRule(context.Background(), api.PostApiAgentNetworkBudgetRulesJSONRequestBody{ + Name: "Org monthly ceiling", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_UpdateBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateBudgetRule(context.Background(), "ainbud_test", api.PutApiAgentNetworkBudgetRulesRuleIdJSONRequestBody{ + Name: "Org monthly ceiling", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_DeleteBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteBudgetRule(context.Background(), "ainbud_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_GetSettings_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkSettings) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkSettings, *ret) + }) +} + +// TestAgentNetwork_GetSettings_UnbootstrappedDefaults pins the settings-read +// contract: an unbootstrapped account answers 200 with the defaults and empty +// cluster/subdomain/endpoint, which the client passes through untouched. +func TestAgentNetwork_GetSettings_UnbootstrappedDefaults(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(api.AgentNetworkSettings{ + EnableLogCollection: true, + AccessLogRetentionDays: ptr(30), + }) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.NoError(t, err) + assert.Empty(t, ret.Endpoint, "empty endpoint is the not-bootstrapped signal") + assert.True(t, ret.EnableLogCollection, "defaults must pass through") + }) +} + +func TestAgentNetwork_GetSettings_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "no", Code: 403}) + w.WriteHeader(403) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.GetSettings(context.Background()) + require.Error(t, err) + assert.Equal(t, "no", err.Error()) + }) +} + +// TestAgentNetwork_GetSettings_LegacyNullBody pins the compatibility shim for +// management servers that answered 200 with a JSON null body before the +// defaults contract: the client translates that shape into an IsNotFound +// error instead of returning a bogus zero-valued settings object or +// fabricating defaults the server never stated. +func TestAgentNetwork_GetSettings_LegacyNullBody(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + _, err := w.Write([]byte("null")) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.Error(t, err) + assert.Nil(t, ret) + assert.True(t, rest.IsNotFound(err), "the legacy 200+null shape must surface as IsNotFound") + }) +} + +func TestAgentNetwork_UpdateSettings_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + var req api.PutApiAgentNetworkSettingsJSONRequestBody + require.NoError(t, json.Unmarshal(reqBytes, &req)) + require.NotNil(t, req.Cluster, "bootstrap cluster must be on the wire") + assert.Equal(t, "eu.proxy.netbird.io", *req.Cluster) + assert.True(t, req.EnableLogCollection) + retBytes, _ := json.Marshal(testAgentNetworkSettings) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{ + Cluster: ptr("eu.proxy.netbird.io"), + EnableLogCollection: true, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkSettings, *ret) + }) +} + +func TestAgentNetwork_UpdateSettings_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "cluster is immutable once assigned (current: eu.proxy.netbird.io)", Code: 422}) + w.WriteHeader(422) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{ + Cluster: ptr("us.proxy.netbird.io"), + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "immutable") + }) +} diff --git a/shared/management/client/rest/client.go b/shared/management/client/rest/client.go index 43312b9e6..6154a6637 100644 --- a/shared/management/client/rest/client.go +++ b/shared/management/client/rest/client.go @@ -147,6 +147,10 @@ type Client struct { // ReverseProxyTokens account-scoped proxy access tokens used to register // self-hosted (bring-your-own-proxy) `netbird proxy` instances. ReverseProxyTokens *ReverseProxyTokensAPI + + // AgentNetwork NetBird Agent Network (AI/LLM gateway) APIs: catalog, + // providers, policies, guardrails, budget rules and account settings. + AgentNetwork *AgentNetworkAPI } // New initialize new Client instance using PAT token @@ -209,6 +213,7 @@ func (c *Client) initialize() { c.ReverseProxyClusters = &ReverseProxyClustersAPI{c} c.ReverseProxyDomains = &ReverseProxyDomainsAPI{c} c.ReverseProxyTokens = &ReverseProxyTokensAPI{c} + c.AgentNetwork = &AgentNetworkAPI{c} } // NewRequest creates and executes new management API request diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 8ad3d932c..551a60e2a 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -5149,12 +5149,12 @@ components: identity_header_user_id: type: string description: | - Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). + Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). example: "x-bf-dim-netbird_user_id" identity_header_groups: type: string description: | - Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. + Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. example: "x-bf-dim-netbird_groups" enabled: type: boolean @@ -5186,6 +5186,8 @@ components: - name - upstream_url - models + - identity_header_user_id + - identity_header_groups - enabled - skip_tls_verification - metadata_disabled @@ -5222,7 +5224,7 @@ components: extra_values: type: object description: | - Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). When present on a request, the whole map replaces the stored values. Empty strings drop the corresponding key. + Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). The request's map replaces the stored values; empty strings drop the corresponding key. additionalProperties: type: string example: @@ -5230,12 +5232,12 @@ components: identity_header_user_id: type: string description: | - Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. When omitted on a request, the stored value is left unchanged; pass an empty string explicitly to clear it (which disables stamping for this dimension). + Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. Empty or omitted disables stamping for this dimension. example: "x-bf-dim-netbird_user_id" identity_header_groups: type: string description: | - Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same omit / empty semantics as `identity_header_user_id`. + Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same semantics as `identity_header_user_id`. example: "x-bf-dim-netbird_groups" enabled: type: boolean @@ -5243,11 +5245,11 @@ components: example: true skip_tls_verification: type: boolean - description: Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. When omitted on update, the stored value is left unchanged. + description: Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. example: false metadata_disabled: type: boolean - description: Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). When omitted on update, the stored value is left unchanged. + description: Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). example: false required: - provider_id @@ -6191,19 +6193,19 @@ components: - cache_cost_usd AgentNetworkSettings: type: object - description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are auto-assigned on first provider create and immutable thereafter. + description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are assigned at bootstrap and immutable thereafter. Before bootstrap the account reads as the default values with empty cluster, subdomain and endpoint. properties: cluster: type: string - description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. + description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. Empty until the account is bootstrapped. example: "eu.proxy.netbird.io" subdomain: type: string - description: Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. + description: Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. Empty until the account is bootstrapped. example: "violet" endpoint: type: string - description: Bare hostname agents call for this account, computed as `.`. + description: Bare hostname agents call for this account, computed as `.`. Empty until the account is bootstrapped. example: "violet.eu.proxy.netbird.io" enable_log_collection: type: boolean @@ -6224,13 +6226,13 @@ components: created_at: type: string format: date-time - description: Timestamp when the settings row was created. + description: Timestamp when the settings row was created. Absent until the account is bootstrapped. readOnly: true example: "2026-04-26T10:30:00Z" updated_at: type: string format: date-time - description: Timestamp when the settings row was last updated. + description: Timestamp when the settings row was last updated. Absent until the account is bootstrapped. readOnly: true example: "2026-04-26T10:30:00Z" required: @@ -6240,12 +6242,14 @@ components: - enable_log_collection - enable_prompt_collection - redact_pii - - created_at - - updated_at AgentNetworkSettingsRequest: type: object - description: Mutable account-level Agent Network settings. Cluster and subdomain are immutable and not accepted here. + description: Account-level Agent Network settings update. The request replaces every mutable field. `cluster` additionally bootstraps the per-account settings row when the account does not have one yet; the subdomain is always server-assigned. properties: + cluster: + type: string + description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. When the account has no settings row yet, providing it bootstraps the row (assigning the subdomain that forms the agent endpoint). The cluster is immutable once assigned — later updates must omit it or send the assigned value; any other value is rejected. + example: "eu.proxy.netbird.io" enable_log_collection: type: boolean description: Whether per-request access-log entries are collected for this account's agent-network traffic. @@ -13690,7 +13694,7 @@ paths: /api/agent-network/settings: get: summary: Retrieve Agent Network settings - description: Returns the per-account Agent Network gateway settings (cluster, subdomain, endpoint). Returns 404 when no provider has been created yet — settings are lazily bootstrapped on first provider create. + description: Returns the per-account Agent Network gateway settings (cluster, subdomain, endpoint). Before the account is bootstrapped — on first provider create (`bootstrap_cluster`) or via PUT with `cluster` — the response carries the default values with empty cluster, subdomain and endpoint. tags: [ Agent Network ] security: - BearerAuth: [ ] @@ -13706,13 +13710,11 @@ paths: "$ref": "#/components/responses/requires_authentication" '403': "$ref": "#/components/responses/forbidden" - '404': - "$ref": "#/components/responses/not_found" '500': "$ref": "#/components/responses/internal_error" put: summary: Update Agent Network settings - description: Updates the mutable account-level Agent Network settings (collection toggles). Cluster and subdomain are immutable and ignored if sent. Returns 404 when settings have not been bootstrapped (no provider created yet). + description: Updates the account-level Agent Network settings; the request replaces every mutable field (collection toggles and retention). When the account has no settings row yet, providing `cluster` bootstraps it (assigning the subdomain that forms the agent endpoint); without `cluster` the request returns 404. Sending a `cluster` different from the assigned one is rejected (the cluster is immutable once assigned). The subdomain is always server-assigned and immutable. tags: [ Agent Network ] security: - BearerAuth: [ ] @@ -13738,6 +13740,8 @@ paths: "$ref": "#/components/responses/forbidden" '404': "$ref": "#/components/responses/not_found" + '422': + "$ref": "#/components/responses/validation_failed" '500': "$ref": "#/components/responses/internal_error" /api/agent-network/budget-rules: diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index 87dd9ccfc..dec644114 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -2275,11 +2275,11 @@ type AgentNetworkProvider struct { // Id Provider ID Id string `json:"id"` - // IdentityHeaderGroups Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. - IdentityHeaderGroups *string `json:"identity_header_groups,omitempty"` + // IdentityHeaderGroups Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. + IdentityHeaderGroups string `json:"identity_header_groups"` - // IdentityHeaderUserId Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). - IdentityHeaderUserId *string `json:"identity_header_user_id,omitempty"` + // IdentityHeaderUserId Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). + IdentityHeaderUserId string `json:"identity_header_user_id"` // MetadataDisabled Whether identity metadata injection is disabled for this provider. When enabled (the default), the proxy stamps the caller's user and authorizing group onto upstream requests as provider-specific metadata (e.g. AWS Bedrock's X-Amzn-Bedrock-Request-Metadata header). Set true to suppress it. MetadataDisabled bool `json:"metadata_disabled"` @@ -2335,16 +2335,16 @@ type AgentNetworkProviderRequest struct { // Enabled Whether the provider is enabled. Defaults to true on create. Enabled *bool `json:"enabled,omitempty"` - // ExtraValues Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). When present on a request, the whole map replaces the stored values. Empty strings drop the corresponding key. + // ExtraValues Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). The request's map replaces the stored values; empty strings drop the corresponding key. ExtraValues *map[string]string `json:"extra_values,omitempty"` - // IdentityHeaderGroups Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same omit / empty semantics as `identity_header_user_id`. + // IdentityHeaderGroups Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same semantics as `identity_header_user_id`. IdentityHeaderGroups *string `json:"identity_header_groups,omitempty"` - // IdentityHeaderUserId Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. When omitted on a request, the stored value is left unchanged; pass an empty string explicitly to clear it (which disables stamping for this dimension). + // IdentityHeaderUserId Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. Empty or omitted disables stamping for this dimension. IdentityHeaderUserId *string `json:"identity_header_user_id,omitempty"` - // MetadataDisabled Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). When omitted on update, the stored value is left unchanged. + // MetadataDisabled Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). MetadataDisabled *bool `json:"metadata_disabled,omitempty"` // Models Models exposed through this endpoint, with the operator's per-1k input/output prices. Empty means all catalog models are allowed at catalog prices. @@ -2356,22 +2356,22 @@ type AgentNetworkProviderRequest struct { // ProviderId Catalog identifier for the upstream AI provider (e.g. openai_api, anthropic_api, azure_openai_api, bedrock_api, vertex_ai_api, mistral_api, custom). ProviderId string `json:"provider_id"` - // SkipTlsVerification Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. When omitted on update, the stored value is left unchanged. + // SkipTlsVerification Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. SkipTlsVerification *bool `json:"skip_tls_verification,omitempty"` // UpstreamUrl Full upstream URL (with scheme) that NetBird forwards traffic to. UpstreamUrl string `json:"upstream_url"` } -// AgentNetworkSettings Per-account Agent Network gateway settings. One row per account; cluster and subdomain are auto-assigned on first provider create and immutable thereafter. +// AgentNetworkSettings Per-account Agent Network gateway settings. One row per account; cluster and subdomain are assigned at bootstrap and immutable thereafter. Before bootstrap the account reads as the default values with empty cluster, subdomain and endpoint. type AgentNetworkSettings struct { // AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. Usage records are retained independently. AccessLogRetentionDays *int `json:"access_log_retention_days,omitempty"` - // Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. + // Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. Empty until the account is bootstrapped. Cluster string `json:"cluster"` - // CreatedAt Timestamp when the settings row was created. + // CreatedAt Timestamp when the settings row was created. Absent until the account is bootstrapped. CreatedAt *time.Time `json:"created_at,omitempty"` // EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic. @@ -2380,24 +2380,27 @@ type AgentNetworkSettings struct { // EnablePromptCollection Master switch for request/response prompt capture. Capture runs only when this is on AND a policy guardrail also enables it. EnablePromptCollection bool `json:"enable_prompt_collection"` - // Endpoint Bare hostname agents call for this account, computed as `.`. + // Endpoint Bare hostname agents call for this account, computed as `.`. Empty until the account is bootstrapped. Endpoint string `json:"endpoint"` // RedactPii Whether captured prompts have PII redacted. Effective redaction is the OR of this and any policy guardrail's redact setting. RedactPii bool `json:"redact_pii"` - // Subdomain Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. + // Subdomain Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. Empty until the account is bootstrapped. Subdomain string `json:"subdomain"` - // UpdatedAt Timestamp when the settings row was last updated. + // UpdatedAt Timestamp when the settings row was last updated. Absent until the account is bootstrapped. UpdatedAt *time.Time `json:"updated_at,omitempty"` } -// AgentNetworkSettingsRequest Mutable account-level Agent Network settings. Cluster and subdomain are immutable and not accepted here. +// AgentNetworkSettingsRequest Account-level Agent Network settings update. The request replaces every mutable field. `cluster` additionally bootstraps the per-account settings row when the account does not have one yet; the subdomain is always server-assigned. type AgentNetworkSettingsRequest struct { // AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. AccessLogRetentionDays *int `json:"access_log_retention_days,omitempty"` + // Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. When the account has no settings row yet, providing it bootstraps the row (assigning the subdomain that forms the agent endpoint). The cluster is immutable once assigned — later updates must omit it or send the assigned value; any other value is rejected. + Cluster *string `json:"cluster,omitempty"` + // EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic. EnableLogCollection bool `json:"enable_log_collection"` diff --git a/shared/management/networkmap/decode.go b/shared/management/networkmap/decode.go index d15117b6e..4864a9dff 100644 --- a/shared/management/networkmap/decode.go +++ b/shared/management/networkmap/decode.go @@ -228,15 +228,17 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents, return c, nil } +// decodeAccountNetwork never returns nil — Calculate() dereferences +// c.Network unconditionally, and servers that predate the fix omit the field +// entirely from the empty-components envelope. func decodeAccountNetwork(an *proto.AccountNetwork) *types.Network { + n := &types.Network{} if an == nil { - return nil - } - n := &types.Network{ - Identifier: an.Identifier, - Dns: an.Dns, - Serial: an.Serial, + return n } + n.Identifier = an.Identifier + n.Dns = an.Dns + n.Serial = an.Serial if an.NetCidr != "" { if _, ipnet, err := net.ParseCIDR(an.NetCidr); err == nil && ipnet != nil { n.Net = *ipnet diff --git a/shared/management/networkmap/envelope_test.go b/shared/management/networkmap/envelope_test.go index a81478aff..92e1916da 100644 --- a/shared/management/networkmap/envelope_test.go +++ b/shared/management/networkmap/envelope_test.go @@ -221,6 +221,66 @@ func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) { "client-side Calculate must connect the same remote peers as the server") } +// TestEnvelopeToNetworkMap_EmptyComponents covers the graceful-degrade path +// the server takes for a peer that is missing from the account or absent from +// the validated-peers map. The legacy server short-circuited before +// Calculate() and shipped a NetworkMap carrying only the account Network; the +// components path runs Calculate() on the client instead, so the envelope must +// carry Network or the client panics dereferencing a nil *types.Network. +func TestEnvelopeToNetworkMap_EmptyComponents(t *testing.T) { + localPeerKey := randomWgKey(t) + c := types.EmptyNetworkMapComponents(&types.NetworkMapComponents{ + PeerID: "peer-A", + Network: &types.Network{ + Identifier: "net-empty", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 7, + }, + Peers: map[string]*types.ComponentPeer{ + "peer-A": {ID: "peer-A", Key: localPeerKey, IP: netip.AddrFrom4([4]byte{100, 64, 0, 1})}, + }, + }) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + require.NotNil(t, envelope.GetFull().Network, "empty envelope must carry the account Network") + + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + require.NoError(t, err, "EnvelopeToNetworkMap must degrade gracefully on empty components") + require.Equal(t, uint64(7), result.NetworkMap.Serial) + require.Empty(t, result.NetworkMap.RemotePeers, "unvalidated peer connects to nobody") +} + +// TestEnvelopeToNetworkMap_MissingNetwork simulates a server that omits +// AccountNetwork from the envelope. Clients must degrade rather than panic, so +// they survive talking to a management server that predates the encoder fix. +func TestEnvelopeToNetworkMap_MissingNetwork(t *testing.T) { + c, localPeerKey := buildSmokeComponents(t) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + envelope.GetFull().Network = nil + + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + require.NoError(t, err, "a missing AccountNetwork must not panic the client") + require.NotNil(t, result.Components.Network) + require.NotEmpty(t, result.NetworkMap.RemotePeers, "the rest of the snapshot stays usable") +} + // buildSmokeComponents returns a minimal NetworkMapComponents (2 peers, 1 // group, 1 allow policy) plus the receiving peer's WG public key. Sufficient // to validate the encode → marshal → decode → Calculate pipeline produces diff --git a/shared/relay/client/guard.go b/shared/relay/client/guard.go index 98b1b333e..d18534d9d 100644 --- a/shared/relay/client/guard.go +++ b/shared/relay/client/guard.go @@ -156,9 +156,11 @@ func (g *Guard) notifyReconnected() { func (g *Guard) exponentTicker(ctx context.Context) *backoff.Ticker { bo := backoff.WithContext(&backoff.ExponentialBackOff{ InitialInterval: 2 * time.Second, - Multiplier: 2, - MaxInterval: g.maxBackoffInterval, - Clock: backoff.SystemClock, + // Spreads the reconnects of every client that lost the same relay server. + RandomizationFactor: backoff.DefaultRandomizationFactor, + Multiplier: 2, + MaxInterval: g.maxBackoffInterval, + Clock: backoff.SystemClock, }, ctx) return backoff.NewTicker(bo)