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/.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/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..4ac006795 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,514 @@ +# 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 + +- [NetBird Agent Guidelines](#netbird-agent-guidelines) + - [Contents](#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) + - [Repo-wide principles](#repo-wide-principles) + - [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/` | + +## Repo-wide principles + +1. **Run `go fmt` on every modified Go file.** Formatting is not optional. +2. **Zero unaddressed diagnostics.** Fix IDE and linter warnings 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 adds + shared state. Guard maps and slices with a mutex, keep critical sections + short, and run `go test -race` on the touched packages. +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. + +## 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. + +## 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 + +- **90 characters per line.** Wrap the comment, do not run past it. +- **250 characters per comment**, roughly three wrapped lines. Doc comments on + exported identifiers may exceed it when the API genuinely needs the + explanation; inline comments inside a function body may not. + +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. + +```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. + +- **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. + +- **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 dbcc59d79..ffe4ba78e 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 6fa027cae..f22ac908d 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -582,12 +582,7 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) } e.stateManager.Start() - initialRoutes, dnsConfig, dnsFeatureFlag, err := e.readInitialSettings() - if err != nil { - return fmt.Errorf("read initial settings: %w", err) - } - - dnsServer, err := e.newDnsServer(dnsConfig) + dnsServer, err := e.newDnsServer() if err != nil { return fmt.Errorf("create dns server: %w", err) } @@ -605,10 +600,8 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) WGInterface: e.wgInterface, StatusRecorder: e.statusRecorder, RelayManager: e.relayManager, - InitialRoutes: initialRoutes, StateManager: e.stateManager, DNSServer: dnsServer, - DNSFeatureFlag: dnsFeatureFlag, PeerStore: e.peerStore, DisableClientRoutes: e.config.DisableClientRoutes, DisableServerRoutes: e.config.DisableServerRoutes, @@ -2122,43 +2115,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.ServerVNCAllowed, - 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 { @@ -2193,7 +2149,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() { @@ -2206,7 +2162,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 @@ -2218,7 +2174,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 9afdf26f1..a5521636a 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -139,6 +139,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 { @@ -374,7 +381,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 { @@ -629,7 +663,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 } @@ -690,7 +736,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 } @@ -845,7 +891,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 } @@ -1808,6 +1854,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/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/main.go b/client/ui/main.go index dfa1a39c5..cd09dd0e8 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 @@ -210,7 +209,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..49acb7917 100644 --- a/client/ui/preferences/store.go +++ b/client/ui/preferences/store.go @@ -246,6 +246,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/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/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..3093c693b 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -44,7 +44,7 @@ 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 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/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 263e090c3..47cc81776 100644 --- a/go.mod +++ b/go.mod @@ -117,7 +117,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 @@ -342,3 +342,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 f773fc6c8..7d1ef4b4c 100644 --- a/go.sum +++ b/go.sum @@ -497,6 +497,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= @@ -667,8 +669,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..87b0d65ee 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)") @@ -247,6 +271,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 +281,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 +307,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 +330,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 +407,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 +516,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 +600,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 +651,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 +679,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: "" @@ -632,6 +702,16 @@ server: 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/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 1c78af9d0..2a1e521aa 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -24,13 +24,13 @@ import ( "github.com/netbirdio/netbird/encryption" "github.com/netbirdio/netbird/formatter/hook" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/activity" activitystore "github.com/netbirdio/netbird/management/server/activity/store" - "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" nbcache "github.com/netbirdio/netbird/management/server/cache" nbContext "github.com/netbirdio/netbird/management/server/context" nbhttp "github.com/netbirdio/netbird/management/server/http" @@ -184,6 +184,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 { @@ -215,6 +219,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 62dbef315..e0ac2c5a9 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 a966dbe8a..89aabe608 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -1685,14 +1685,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 7f10b603f..b48ecf98f 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 1d224adf8..5c27a09b1 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -5162,12 +5162,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 @@ -5199,6 +5199,8 @@ components: - name - upstream_url - models + - identity_header_user_id + - identity_header_groups - enabled - skip_tls_verification - metadata_disabled @@ -5235,7 +5237,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: @@ -5243,12 +5245,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 @@ -5256,11 +5258,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 @@ -6204,19 +6206,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 @@ -6237,13 +6239,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: @@ -6253,12 +6255,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. @@ -13703,7 +13707,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: [ ] @@ -13719,13 +13723,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: [ ] @@ -13751,6 +13753,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 ed727ab5d..8aab80f81 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 b0858f11b..f7b4ca027 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 e1469d0b2..ae0df9759 100644 --- a/shared/management/networkmap/envelope_test.go +++ b/shared/management/networkmap/envelope_test.go @@ -292,6 +292,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)