From d681670a9d17b45f407e487896973c03cd3a6756 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Mon, 27 Jul 2026 19:07:58 +0900 Subject: [PATCH 01/32] [misc] Restore the rootless-latest docker tag (#6914) --- .goreleaser.yaml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.goreleaser.yaml b/.goreleaser.yaml index 0753b1012..8dd05a192 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -273,8 +273,8 @@ dockers_v2: - netbirdio/netbird - ghcr.io/netbirdio/netbird tags: - - "v{{ .Version }}-rootless" - - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}" + - "{{ .Version }}-rootless" + - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}rootless-latest{{ end }}" dockerfile: client/Dockerfile-rootless extra_files: - client/netbird-entrypoint.sh From aa13928b76e0a1741bc23875fa0833347593729a Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 27 Jul 2026 15:53:09 +0200 Subject: [PATCH 02/32] [client] Export agent version info for iOS (#6918) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Export agent version info for iOS ## Issue ticket number and link ## Stack ### Checklist - [ ] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [x] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. --- client/ios/NetBirdSDK/version.go | 12 ++++++++++++ 1 file changed, 12 insertions(+) create mode 100644 client/ios/NetBirdSDK/version.go diff --git a/client/ios/NetBirdSDK/version.go b/client/ios/NetBirdSDK/version.go new file mode 100644 index 000000000..606ad18e2 --- /dev/null +++ b/client/ios/NetBirdSDK/version.go @@ -0,0 +1,12 @@ +//go:build ios + +package NetBirdSDK + +import "github.com/netbirdio/netbird/version" + +// GoClientVersion returns the NetBird Go client version that was baked into +// the framework at compile time via +// -ldflags "-X github.com/netbirdio/netbird/version.version=". +func GoClientVersion() string { + return version.NetbirdVersion() +} From 1816a020c46847fa75437b1e915134673ff33b16 Mon Sep 17 00:00:00 2001 From: Misha Bragin Date: Mon, 27 Jul 2026 16:15:36 +0200 Subject: [PATCH 03/32] [management, proxy] Add Claude Opus 5 (#6895) --- proxy/internal/llm/pricing/defaults_coverage_test.go | 4 ++-- proxy/internal/llm/pricing/defaults_pricing.yaml | 10 ++++++++++ 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/proxy/internal/llm/pricing/defaults_coverage_test.go b/proxy/internal/llm/pricing/defaults_coverage_test.go index be23682da..8df1557ea 100644 --- a/proxy/internal/llm/pricing/defaults_coverage_test.go +++ b/proxy/internal/llm/pricing/defaults_coverage_test.go @@ -28,12 +28,12 @@ func TestDefaultTable_FirstPartyModelCoverage(t *testing.T) { "ministral-8b-latest", "mistral-embed", }, "anthropic": { - "claude-fable-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6", + "claude-fable-5", "claude-opus-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6", "claude-opus-4-1", "claude-sonnet-4-6", "claude-sonnet-4-5", "claude-haiku-4-5", }, // bedrock keys are the normalized ids the request parser emits. "bedrock": { - "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6", + "anthropic.claude-opus-5", "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6", "anthropic.claude-opus-4-1", "anthropic.claude-sonnet-4-6", "anthropic.claude-sonnet-4-5", "anthropic.claude-haiku-4-5", "meta.llama3-3-70b-instruct", "amazon.nova-pro", "amazon.nova-lite", "amazon.nova-micro", "amazon.nova-2-lite", diff --git a/proxy/internal/llm/pricing/defaults_pricing.yaml b/proxy/internal/llm/pricing/defaults_pricing.yaml index 3fba8fe3f..988426105 100644 --- a/proxy/internal/llm/pricing/defaults_pricing.yaml +++ b/proxy/internal/llm/pricing/defaults_pricing.yaml @@ -180,6 +180,11 @@ anthropic: output_per_1k: 0.050 cache_read_per_1k: 0.001 cache_creation_per_1k: 0.0125 + claude-opus-5: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 claude-opus-4-8: input_per_1k: 0.005 output_per_1k: 0.025 @@ -236,6 +241,11 @@ bedrock: # eu.anthropic.claude-sonnet-4-5-20250929-v1:0 -> anthropic.claude-sonnet-4-5. # Anthropic-on-Bedrock keeps the additive cache buckets (read ≈0.1x input, # write ≈1.25x input); Nova / Llama report no cache, so cost is input+output. + anthropic.claude-opus-5: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 anthropic.claude-opus-4-8: input_per_1k: 0.005 output_per_1k: 0.025 From 9b4a5df9250b81eedfff899e76e026896e7158fa Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 27 Jul 2026 20:08:35 +0200 Subject: [PATCH 04/32] [client] Use platform installer URL for manual update downloads (#6922) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Wails UI regressed the non-enforced update download to the plain GitHub releases page. Restore the old Fyne behavior: the tray update item and the About card's Get installer button now open version.DownloadUrl(), which points to the direct installer download (pkgs.netbird.io) per OS/arch and falls back to the generic install page where no installer exists. Also fix the dead brew detection on darwin: exec.Command passed the whole pipeline as a single argument to brew, so the check always failed. Query the netbird formula and netbird-ui cask explicitly via exit codes instead. ## Describe your changes ## Issue ticket number and link ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **New Features** * Added a direct installer download option for manual application updates. * Non-enforced updates now open the appropriate platform-specific installer instead of the general releases page. * **Bug Fixes** * Improved macOS download behavior for Homebrew installations. * Preserved architecture-specific downloads for Intel and Apple silicon Macs. * Updated the tray “About” links to use the GitHub repository and documentation instead of the releases page. --- .../src/modules/auto-update/UpdateVersionCard.tsx | 13 ++++++++----- client/ui/services/update.go | 7 +++++++ client/ui/tray.go | 5 ++--- client/ui/tray_update.go | 7 ++++--- 4 files changed, 21 insertions(+), 11 deletions(-) diff --git a/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx b/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx index 4fb5e2586..4861c492a 100644 --- a/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx +++ b/client/ui/frontend/src/modules/auto-update/UpdateVersionCard.tsx @@ -2,6 +2,7 @@ import { type ReactNode } from "react"; import { useTranslation } from "react-i18next"; import { Browser } from "@wailsio/runtime"; import { DownloadIcon, NotepadText } from "lucide-react"; +import { Update as UpdateSvc } from "@bindings/services"; import { Button } from "@/components/buttons/Button"; import { useClientVersion } from "@/contexts/ClientVersionContext"; import { cn } from "@/lib/cn"; @@ -14,6 +15,12 @@ function openUrl(url: string) { }); } +function openInstallerDownload() { + UpdateSvc.DownloadURL() + .then(openUrl) + .catch(() => openUrl(GITHUB_RELEASES)); +} + export function UpdateVersionCard() { const { t } = useTranslation(); const { updateVersion, enforced, triggerUpdate } = useClientVersion(); @@ -37,11 +44,7 @@ export function UpdateVersionCard() { {t("update.card.installNow")} ) : ( - diff --git a/client/ui/services/update.go b/client/ui/services/update.go index 753177d45..b743b9858 100644 --- a/client/ui/services/update.go +++ b/client/ui/services/update.go @@ -10,6 +10,7 @@ import ( "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/ui/updater" + "github.com/netbirdio/netbird/version" ) // UpdateResult mirrors TriggerUpdateResponse. @@ -33,6 +34,12 @@ func (s *Update) GetState() updater.State { return s.holder.Get() } +// DownloadURL returns the platform-appropriate installer download link for +// manual (non-enforced) updates. +func (s *Update) DownloadURL() string { + return version.DownloadUrl() +} + // Quit exits the app. Scheduled off the calling goroutine so the JS caller's // response returns before the runtime tears down. func (s *Update) Quit() { diff --git a/client/ui/tray.go b/client/ui/tray.go index 63b6a46ec..49e883f89 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -32,9 +32,8 @@ const ( quitDownTimeout = 5 * time.Second - urlGitHubRepo = "https://github.com/netbirdio/netbird" - urlGitHubReleases = "https://github.com/netbirdio/netbird/releases/latest" - urlDocs = "https://docs.netbird.io" + urlGitHubRepo = "https://github.com/netbirdio/netbird" + urlDocs = "https://docs.netbird.io" ) // TrayServices bundles the services the tray menu needs, grouped so NewTray diff --git a/client/ui/tray_update.go b/client/ui/tray_update.go index 2d79cff05..1a377dfa3 100644 --- a/client/ui/tray_update.go +++ b/client/ui/tray_update.go @@ -13,6 +13,7 @@ import ( "github.com/netbirdio/netbird/client/ui/services" "github.com/netbirdio/netbird/client/ui/updater" + "github.com/netbirdio/netbird/version" ) // trayUpdater owns the tray UI that reacts to auto-update. Composed inside Tray. @@ -76,15 +77,15 @@ func (u *trayUpdater) applyLanguage() { u.refreshMenuItem(state) } -// handleClick opens the GitHub releases page when not Enforced, otherwise shows -// the progress page and asks the daemon to start the installer. +// handleClick opens the installer download link when not Enforced, otherwise +// shows the progress page and asks the daemon to start the installer. func (u *trayUpdater) handleClick() { u.mu.Lock() state := u.state u.mu.Unlock() if !state.Enforced { - _ = u.app.Browser.OpenURL(urlGitHubReleases) + _ = u.app.Browser.OpenURL(version.DownloadUrl()) return } From bab5572a749d86adb8b9899d077fabd3b910d91a Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Tue, 28 Jul 2026 03:43:59 +0900 Subject: [PATCH 05/32] [management, proxy] scope agent-network model allowlist per policy/group and provider (#6905) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes The model-allowlist guardrail was merged into one account-wide union and enforced flat on every request, ignoring which policy/group/provider authorised it. With multiple policies — especially a mix of guardrailed and un-guardrailed ones — this caused: - **false-allow**: a model allowlisted for one group/provider leaked to any caller; and - **false-deny**: an un-guardrailed policy (intended unrestricted) was blocked by another policy's allowlist. Enforcement is now policy/group-aware, mirroring `llm_limit_check`: - **Management (`SelectPolicyForRequest`) is authoritative.** It uses the request model (already carried in `CheckLLMPolicyLimitsRequest.model`, previously ignored) to keep only applicable policies whose guardrails permit the model; no allowlist-enabled guardrail = unrestricted. Denies `llm_policy.model_blocked` when policies govern the (provider, groups) but none permits the model. - **Proxy `llm_guardrail` becomes a per-provider fail-closed backstop.** The synthesiser emits an allowlist only for providers every authorising policy restricts; the middleware keys off the resolved provider id and keeps unknown-model fail-closed. --- docs/agent-networks/01-end-to-end-flows.md | 17 +- .../modules/21-management-agentnetwork.md | 2 +- .../modules/31-proxy-middleware-builtin.md | 2 +- e2e/agentnetwork/guardrail_block_test.go | 209 +++++++++ .../guardrail_groupswitch_test.go | 205 +++++++++ .../guardrail_multipolicy_test.go | 201 +++++++++ .../guardrail_pergroup_providers_test.go | 422 ++++++++++++++++++ e2e/harness/proxy.go | 12 +- .../internals/modules/agentnetwork/manager.go | 7 +- .../modules/agentnetwork/policyselect.go | 108 +++++ .../agentnetwork/policyselect_model_test.go | 329 ++++++++++++++ .../modules/agentnetwork/synthesizer.go | 182 ++++---- .../synthesizer_guardrail_realstore_test.go | 6 +- .../synthesizer_provider_allowlist_test.go | 95 ++++ .../modules/agentnetwork/synthesizer_test.go | 8 +- management/internals/shared/grpc/proxy.go | 1 + .../grpc/proxy_llm_policy_limits_test.go | 138 ++++++ proxy/internal/auth/tunnel_cache.go | 39 +- proxy/internal/auth/tunnel_cache_test.go | 29 ++ .../builtin/llm_guardrail/factory.go | 37 +- .../builtin/llm_guardrail/middleware.go | 42 +- .../builtin/llm_guardrail/middleware_test.go | 143 +++++- .../builtin/llm_limit_check/middleware.go | 17 +- .../llm_limit_check/middleware_test.go | 40 ++ .../guardrail_allowlist_test.go | 11 +- 25 files changed, 2148 insertions(+), 154 deletions(-) create mode 100644 e2e/agentnetwork/guardrail_block_test.go create mode 100644 e2e/agentnetwork/guardrail_groupswitch_test.go create mode 100644 e2e/agentnetwork/guardrail_multipolicy_test.go create mode 100644 e2e/agentnetwork/guardrail_pergroup_providers_test.go create mode 100644 management/internals/modules/agentnetwork/policyselect_model_test.go create mode 100644 management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go create mode 100644 management/internals/shared/grpc/proxy_llm_policy_limits_test.go diff --git a/docs/agent-networks/01-end-to-end-flows.md b/docs/agent-networks/01-end-to-end-flows.md index 7264f3768..b8891001b 100644 --- a/docs/agent-networks/01-end-to-end-flows.md +++ b/docs/agent-networks/01-end-to-end-flows.md @@ -109,7 +109,7 @@ sequenceDiagram Chk->>Inj: continue Inj->>Inj: inject NetBird identity headers per provider config Inj->>Grd: continue - Grd->>Grd: enforce model allowlist + Grd->>Grd: enforce per-provider allowlist (fail-closed backstop) Grd->>Up: forward (over WireGuard) Up-->>Resp: response (JSON or SSE stream) Resp->>Resp: parse usage tokens, completion @@ -135,6 +135,21 @@ sequenceDiagram (`redact_pii = settings.RedactPii`). Phones, emails, credit cards, PII names — see `redact.go` for the full set. See [`modules/31-proxy-middleware-builtin.md`](modules/31-proxy-middleware-builtin.md). +- The model allowlist is enforced in TWO places. `CheckLLMPolicyLimits` + is authoritative: it resolves the policy that governs this + (provider, caller-groups) and denies (`deny_code = llm_policy.model_blocked`) + when no applicable policy permits the model — so an allowlist scoped to + one group/provider never leaks to another, and an un-guardrailed policy + is genuinely unrestricted. `llm_guardrail` is a per-provider fail-closed + backstop: it only carries an allowlist for a provider every authorising + policy restricts, and blocks unknown/undetermined models even when + management is unreachable. Because that backstop allowlist is the UNION + of every restricting policy's models, per-group narrowing lives only in + the authoritative check: during a `CheckLLMPolicyLimits` outage + `llm_limit_check` fails open, so a caller can reach any model in the + provider's union — a group scoped to model A could reach model B if + another group restricts the same provider to B. This is the documented + fail-open trade-off; a future flag may switch it to fail-closed. - SSE streaming requires special handling on the response side; the parser must handle partial chunks without buffering the whole stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.md). diff --git a/docs/agent-networks/modules/21-management-agentnetwork.md b/docs/agent-networks/modules/21-management-agentnetwork.md index b64c1ba20..cc74206e9 100644 --- a/docs/agent-networks/modules/21-management-agentnetwork.md +++ b/docs/agent-networks/modules/21-management-agentnetwork.md @@ -122,7 +122,7 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest | on_request | 1 | `llm_router` | `{"providers":[{id, models[], upstream_*, auth_header_*, allowed_group_ids[]}]}` | **true** | | on_request | 2 | `llm_limit_check` | `{}` | – | | on_request | 3 | `llm_identity_inject` | `{"providers":[{provider_id, header_pair?, json_metadata?, extra_headers?}]}` | **true** | - | on_request | 4 | `llm_guardrail` | `{"model_allowlist"?, "prompt_capture":{enabled,redact_pii}}` | – | + | on_request | 4 | `llm_guardrail` | `{"provider_allowlists"?: {providerID: []model}, "prompt_capture":{enabled,redact_pii}}` | – | | on_response | 5 | `llm_limit_record` | `{}` (runs LAST at runtime) | – | | on_response | 6 | `cost_meter` | `{}` | – | | on_response | 7 | `llm_response_parser` | `{"capture_completion": , "redact_pii"?: true}` | – | diff --git a/docs/agent-networks/modules/31-proxy-middleware-builtin.md b/docs/agent-networks/modules/31-proxy-middleware-builtin.md index 904de6424..efe1bc4ce 100644 --- a/docs/agent-networks/modules/31-proxy-middleware-builtin.md +++ b/docs/agent-networks/modules/31-proxy-middleware-builtin.md @@ -244,7 +244,7 @@ no mocks. Tests: `TestChain_AllowPath_StampsAttributionAndRecordsCounter` | `llm_router` | `{providers: [{id, models, upstream_scheme, upstream_host, upstream_path?, auth_header_name, auth_header_value, allowed_group_ids}]}` | | `llm_limit_check` | `{}` — pulls `MgmtClient` from `FactoryContext` | | `llm_identity_inject` | `{providers: [{provider_id, header_pair?|json_metadata?, extra_headers?}]}` | -| `llm_guardrail` | `{model_allowlist: []string, prompt_capture: {enabled, redact_pii}}` | +| `llm_guardrail` | `{provider_allowlists: {providerID: []string}, prompt_capture: {enabled, redact_pii}}` — allowlist keyed by resolved provider id; a provider absent from the map is unrestricted (fail-closed backstop; authoritative per-policy/group check is management's `CheckLLMPolicyLimits`) | | `llm_response_parser` | `{redact_pii?, capture_completion?: *bool}` | | `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) | | `llm_limit_record` | `{}` — same pattern as `llm_limit_check` | diff --git a/e2e/agentnetwork/guardrail_block_test.go b/e2e/agentnetwork/guardrail_block_test.go new file mode 100644 index 000000000..c4b22ae25 --- /dev/null +++ b/e2e/agentnetwork/guardrail_block_test.go @@ -0,0 +1,209 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// pathRoutedGuardrailCase is one provider's self-contained scenario: its own +// provider, its own guardrail whose allowlist holds ONLY that provider's +// allowed model, and its own policy. Each case runs in isolation (its own +// proxy + client), so the guardrail the proxy enforces contains exactly this +// provider's model — never a mixed cross-provider list. +type pathRoutedGuardrailCase struct { + name string + catalogID string // agent-network catalog provider id + wire string // harness.WireVertex | harness.WireBedrock + allowEntry string // the single model id put on the guardrail allowlist + allowModel string // model id sent that MUST be served (200) + blockModel string // model id sent that MUST be denied (403 model_blocked) +} + +// TestGuardrailBlocksUnselectedModel_PathRouted is the end-to-end regression +// guard for the customer report that a model-allowlist guardrail attached to a +// policy has no effect for PATH-ROUTED providers — where the model travels in +// the URL, not the JSON body: Google Vertex (…/models/{model}:rawPredict) and +// AWS Bedrock (/model/{id}/invoke). +// +// Each provider is tested in isolation with a guardrail allowlisting a single +// model of its own: the allowed model (in the URL path) is served (200) and an +// unselected model (in the URL path) is denied 403 by the guardrail +// (llm_policy.model_blocked) before the upstream. The Vertex case mirrors the +// customer verbatim — allow Sonnet, and the unselected model is the exact +// claude-opus-4-6 they reported reaching the model unblocked. The Bedrock case +// sends a region-prefixed, versioned inference-profile id so URL-path model +// normalization is exercised too. +// +// The provider is catch-all (no models), so the router forwards any model and a +// 403 can only come from the guardrail, never model_not_routable. Only the +// upstream LLM is mocked (the vLLM nginx answers any path with 200); management +// synth/reconcile, the proxy middleware chain (URL-path model extraction, +// router, guardrail) and the tunnel are all real, and the guardrail denies +// before the upstream is dialed so the mock cannot influence the block. A +// static bearer api key is used so the router injects a static Authorization +// header instead of minting a GCP token — the only reason path-routed providers +// normally need live credentials — so the test runs with none and is always on. +func TestGuardrailBlocksUnselectedModel_PathRouted(t *testing.T) { + cases := []pathRoutedGuardrailCase{ + { + name: "vertex", + catalogID: "vertex_ai_api", + wire: harness.WireVertex, + allowEntry: "claude-sonnet-4-5", + allowModel: "claude-sonnet-4-5", + blockModel: "claude-opus-4-6", // the customer-reported model + }, + { + name: "bedrock", + catalogID: "bedrock_api", + wire: harness.WireBedrock, + allowEntry: "anthropic.claude-sonnet-4-5", // normalized catalog id + allowModel: "us.anthropic.claude-sonnet-4-5-v1:0", + blockModel: "us.anthropic.claude-opus-4-8-v1:0", + }, + } + + for _, tc := range cases { + tc := tc + t.Run(tc.name, func(t *testing.T) { + runPathRoutedGuardrailCase(t, tc) + }) + } +} + +func runPathRoutedGuardrailCase(t *testing.T, tc pathRoutedGuardrailCase) { + t.Helper() + + const ( + vertexProject = "e2e-project" + vertexRegion = "global" + ) + + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-" + tc.name}) + require.NoError(t, err, "create group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grp.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-guardrail-" + tc.name + "-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grp.Id}, + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + // Catch-all provider (no models) so the router forwards any model; a static + // bearer key means the router injects a static auth header instead of minting + // a GCP token. Bootstraps the cluster if it isn't already. + staticKey := "static-e2e-token" + prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: tc.name, + ProviderId: tc.catalogID, + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + BootstrapCluster: ptr(harness.AgentNetworkCluster), + }) + require.NoError(t, err, "create %s provider", tc.name) + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) }) + + // Guardrail allowlisting ONLY this provider's allowed model. + var gr api.AgentNetworkGuardrailRequest + gr.Name = "e2e-guardrail-" + tc.name + gr.Checks.ModelAllowlist.Enabled = true + gr.Checks.ModelAllowlist.Models = []string{tc.allowEntry} + guard, err := srv.CreateGuardrail(ctx, gr) + require.NoError(t, err, "create guardrail") + t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), guard.Id) }) + + enabled := true + pol, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-" + tc.name, + Enabled: &enabled, + SourceGroups: []string{grp.Id}, + DestinationProviderIds: []string{prov.Id}, + GuardrailIds: &[]string{guard.Id}, + }) + require.NoError(t, err, "create policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), pol.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-"+tc.name+"-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + // Probe first: the GET resolves the endpoint and its first packet wakes the + // lazy proxy peer, so WaitProxyPeer then observes it connected. + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + send := func(model string) (int, string) { + var code int + var body string + var cerr error + switch tc.wire { + case harness.WireVertex: + code, body, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, vertexProject, vertexRegion, model, "Reply with exactly: pong", "") + case harness.WireBedrock: + code, body, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, model, "Reply with exactly: pong", "") + default: + t.Fatalf("unsupported wire %q", tc.wire) + } + require.NoError(t, cerr, "request must reach the proxy for %s", tc.name) + return code, body + } + + // Allowed model (in the URL path) is served. Retry to absorb tunnel/DNS + // jitter on the first call over the freshly warmed tunnel. + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + code, body = send(tc.allowModel) + if code == 200 { + break + } + time.Sleep(5 * time.Second) + } + assert.Equal(t, 200, code, + "allowed %s model (URL path) must be served; body: %s\n=== proxy logs ===\n%s", tc.name, body, px.Logs(context.Background())) + + // Unselected model (in the URL path) must be blocked by the guardrail. + code, body = send(tc.blockModel) + assert.Equal(t, 403, code, + "unselected %s model (URL path) must be denied, not served; body: %s\n=== proxy logs ===\n%s", tc.name, body, px.Logs(context.Background())) + assert.Contains(t, body, "llm_policy.model_blocked", + "%s denial must come from the guardrail allowlist, not routing; body: %s", tc.name, body) +} diff --git a/e2e/agentnetwork/guardrail_groupswitch_test.go b/e2e/agentnetwork/guardrail_groupswitch_test.go new file mode 100644 index 000000000..5f2dce52e --- /dev/null +++ b/e2e/agentnetwork/guardrail_groupswitch_test.go @@ -0,0 +1,205 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// TestGuardrailGroupSwitchTakesEffectAfterTTL proves that moving a peer between +// groups flips its model-allowlist decision once the proxy's tunnel-peer cache +// expires. The peer's groups reach the guardrail via ValidateTunnelPeer, which +// the proxy caches; the switch is invisible until that cache expires. The proxy +// runs with a short NB_PROXY_TUNNEL_CACHE_TTL so the flip happens in seconds +// instead of the 5-minute default. +// +// Setup: one catch-all provider declaring modelA + modelB; polA (grpA -> allow +// modelA) and polB (grpB -> allow modelB). The client starts in grpA. modelA is +// served and modelB denied; after switching the client grpA -> grpB, modelB is +// served and modelA denied. The cross-group deny comes from management's +// per-policy/group CheckLLMPolicyLimits (the proxy backstop carries the union). +func TestGuardrailGroupSwitchTakesEffectAfterTTL(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + modelA = "e2e-model-a" + modelB = "e2e-model-b" + // Short tunnel-cache TTL so a group switch propagates in seconds. + // Exercises the NB_PROXY_TUNNEL_CACHE_TTL override. + cacheTTL = 3 * time.Second + ) + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grpA, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-gswitch-a"}) + require.NoError(t, err, "create group A") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpA.Id) }) + + grpB, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-gswitch-b"}) + require.NoError(t, err, "create group B") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpB.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-gswitch-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grpA.Id}, // client starts in group A + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + staticKey := "static-e2e-token" + prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "gswitch", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: &[]api.AgentNetworkProviderModel{ + {Id: modelA, InputPer1k: 0.001, OutputPer1k: 0.001}, + {Id: modelB, InputPer1k: 0.001, OutputPer1k: 0.001}, + }, + BootstrapCluster: ptr(harness.AgentNetworkCluster), + }) + require.NoError(t, err, "create provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) }) + + mkGuard := func(name, model string) api.AgentNetworkGuardrail { + var gr api.AgentNetworkGuardrailRequest + gr.Name = name + gr.Checks.ModelAllowlist.Enabled = true + gr.Checks.ModelAllowlist.Models = []string{model} + g, gerr := srv.CreateGuardrail(ctx, gr) + require.NoError(t, gerr, "create guardrail %s", name) + t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) }) + return g + } + gA := mkGuard("e2e-gswitch-a", modelA) + gB := mkGuard("e2e-gswitch-b", modelB) + + enabled := true + polA, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-gswitch-a", + Enabled: &enabled, + SourceGroups: []string{grpA.Id}, + DestinationProviderIds: []string{prov.Id}, + GuardrailIds: &[]string{gA.Id}, + }) + require.NoError(t, err, "create policy A") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polA.Id) }) + + polB, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-gswitch-b", + Enabled: &enabled, + SourceGroups: []string{grpB.Id}, + DestinationProviderIds: []string{prov.Id}, + GuardrailIds: &[]string{gB.Id}, + }) + require.NoError(t, err, "create policy B") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polB.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-gswitch-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken, map[string]string{ + "NB_PROXY_TUNNEL_CACHE_TTL": cacheTTL.String(), + }) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + send := func(model string) (int, string) { + code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "") + require.NoError(t, cerr, "request must reach the proxy") + return code, body + } + sendUntil := func(model string, want int, timeout time.Duration) (int, string) { + var code int + var body string + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + code, body = send(model) + if code == want { + return code, body + } + time.Sleep(2 * time.Second) + } + return code, body + } + + // Phase 1 — client is in group A: modelA served, modelB denied. + code, body := sendUntil(modelA, 200, 90*time.Second) + assert.Equal(t, 200, code, + "group-A model must be served while the client is in group A; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + code, body = send(modelB) + assert.Equal(t, 403, code, + "group-B model must be denied while the client is in group A; body: %s", body) + assert.Contains(t, body, "llm_policy.model_blocked") + + // Switch the client peer from group A to group B. + peerID := clientPeerInGroup(t, ctx, grpA.Id) + _, err = srv.API().Groups.Update(ctx, grpB.Id, api.PutApiGroupsGroupIdJSONRequestBody{ + Name: grpB.Name, + Peers: &[]string{peerID}, + }) + require.NoError(t, err, "add peer to group B") + _, err = srv.API().Groups.Update(ctx, grpA.Id, api.PutApiGroupsGroupIdJSONRequestBody{ + Name: grpA.Name, + Peers: &[]string{}, + }) + require.NoError(t, err, "remove peer from group A") + + // Phase 2 — after the short TTL expires the proxy re-validates the peer, + // sees group B, and the decision flips. Poll to absorb TTL + re-validation. + code, body = sendUntil(modelB, 200, 60*time.Second) + assert.Equal(t, 200, code, + "after the group switch + TTL, the group-B model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + code, body = send(modelA) + assert.Equal(t, 403, code, + "after the switch, the old group-A model must be denied; body: %s", body) + assert.Contains(t, body, "llm_policy.model_blocked") +} + +// clientPeerInGroup returns the id of the single peer that is a member of the +// given group — the test client. The proxy peer is never added to test groups. +func clientPeerInGroup(t *testing.T, ctx context.Context, groupID string) string { + t.Helper() + peers, err := srv.API().Peers.List(ctx) + require.NoError(t, err, "list peers") + for _, p := range peers { + for _, g := range p.Groups { + if g.Id == groupID { + return p.Id + } + } + } + t.Fatalf("no peer found in group %s", groupID) + return "" +} diff --git a/e2e/agentnetwork/guardrail_multipolicy_test.go b/e2e/agentnetwork/guardrail_multipolicy_test.go new file mode 100644 index 000000000..1d664fb58 --- /dev/null +++ b/e2e/agentnetwork/guardrail_multipolicy_test.go @@ -0,0 +1,201 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// TestGuardrailMultiPolicyModelAllowlist: modelSelected served (200), grpOther's +// modelOther denied for the grpMain client (403 model_blocked, no cross-group +// leak), and openModel on the un-guardrailed policy's provider served (200). +func TestGuardrailMultiPolicyModelAllowlist(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + modelSelected = "e2e-selected" + modelOther = "e2e-other" + openModel = "e2e-open" + ) + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-main"}) + require.NoError(t, err, "create main group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) }) + + grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-other"}) + require.NoError(t, err, "create other group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-guardrail-mp-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grpMain.Id}, // client joins grpMain only + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + staticKey := "static-e2e-token" + models := func(ids ...string) *[]api.AgentNetworkProviderModel { + out := make([]api.AgentNetworkProviderModel, 0, len(ids)) + for _, id := range ids { + out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001}) + } + return &out + } + + // pRestricted declares the two guardrailed models so routing is deterministic + // (model -> provider). Created first, so it carries the bootstrap cluster. + pRestricted, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "restricted", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: models(modelSelected, modelOther), + BootstrapCluster: ptr(harness.AgentNetworkCluster), + }) + require.NoError(t, err, "create restricted provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pRestricted.Id) }) + + pOpen, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "open", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: models(openModel), + }) + require.NoError(t, err, "create open provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pOpen.Id) }) + + mkGuardrail := func(name, model string) api.AgentNetworkGuardrail { + var gr api.AgentNetworkGuardrailRequest + gr.Name = name + gr.Checks.ModelAllowlist.Enabled = true + gr.Checks.ModelAllowlist.Models = []string{model} + g, gerr := srv.CreateGuardrail(ctx, gr) + require.NoError(t, gerr, "create guardrail %s", name) + t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) }) + return g + } + gMain := mkGuardrail("e2e-guardrail-mp-main", modelSelected) + gOther := mkGuardrail("e2e-guardrail-mp-other", modelOther) + + enabled := true + // polMain: grpMain restricted to modelSelected on pRestricted. + polMain, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-mp-main", + Enabled: &enabled, + SourceGroups: []string{grpMain.Id}, + DestinationProviderIds: []string{pRestricted.Id}, + GuardrailIds: &[]string{gMain.Id}, + }) + require.NoError(t, err, "create main policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) }) + + // polOther: grpOther restricted to modelOther on the SAME provider. The + // client is not in grpOther, so modelOther must never be usable by it. + polOther, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-mp-other", + Enabled: &enabled, + SourceGroups: []string{grpOther.Id}, + DestinationProviderIds: []string{pRestricted.Id}, + GuardrailIds: &[]string{gOther.Id}, + }) + require.NoError(t, err, "create other policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) }) + + // polOpen: grpMain on pOpen with NO guardrail — unrestricted. + polOpen, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-guardrail-mp-open", + Enabled: &enabled, + SourceGroups: []string{grpMain.Id}, + DestinationProviderIds: []string{pOpen.Id}, + }) + require.NoError(t, err, "create open policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOpen.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-mp-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + send := func(model string) (int, string) { + code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "") + require.NoError(t, cerr, "request must reach the proxy") + return code, body + } + // sendUntil200 absorbs first-call tunnel/DNS jitter on the freshly warmed tunnel. + sendUntil200 := func(model string) (int, string) { + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + code, body = send(model) + if code == 200 { + break + } + time.Sleep(5 * time.Second) + } + return code, body + } + + t.Run("selected model allowed for its group", func(t *testing.T) { + code, body := sendUntil200(modelSelected) + assert.Equal(t, 200, code, + "grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + }) + + t.Run("other group's model does not leak", func(t *testing.T) { + // modelOther is allowlisted only for grpOther. The grpMain client must be + // denied by management's per-policy/group check — not waved through by an + // account-wide union. This is the security-critical wrong-ALLOW guard. + code, body := send(modelOther) + assert.Equal(t, 403, code, + "another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + assert.Contains(t, body, "llm_policy.model_blocked", + "denial must be a model-allowlist decision; body: %s", body) + }) + + t.Run("unguarded policy leaves its provider unrestricted", func(t *testing.T) { + // polOpen carries no guardrail, so pOpen is unrestricted for grpMain. The + // old account-wide union would have blocked openModel (it is on no + // allowlist); it must now be served — the false-DENY guard. + code, body := sendUntil200(openModel) + assert.Equal(t, 200, code, + "an un-guardrailed policy's provider must not be blocked by another policy's allowlist; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + }) +} diff --git a/e2e/agentnetwork/guardrail_pergroup_providers_test.go b/e2e/agentnetwork/guardrail_pergroup_providers_test.go new file mode 100644 index 000000000..eddae65c3 --- /dev/null +++ b/e2e/agentnetwork/guardrail_pergroup_providers_test.go @@ -0,0 +1,422 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// pergroupCase describes one provider surface for the per-group allowlist matrix. +// selectedReq/otherReq are the model identifiers as they travel in the request +// (URL path for Bedrock/Vertex, body "model" for chat/messages). selectedAllow/ +// otherAllow are the (normalized) forms the guardrail allowlist holds — for +// Bedrock these differ from the request form so path normalization is exercised. +type pergroupCase struct { + name string + catalogID string + wire string // "chat", "messages", "vertex", "bedrock" + models *[]api.AgentNetworkProviderModel + selectedReq string + selectedAllow string + otherReq string + otherAllow string + + providerID string // filled during setup +} + +// TestGuardrailPerGroupAllowlist_AllProviders proves the per-policy/group model +// allowlist end to end across every always-on provider surface, including the +// path-routed ones (Vertex, Bedrock) where the model travels in the URL. +// +// For each provider two policies target it: grpMain (the client) is allowed only +// selectedReq; grpOther (which the client is NOT in) is allowed only otherReq. +// The client must get selectedReq served (200) and otherReq denied (403, +// llm_policy.model_blocked) — the cross-group no-leak property. The deny is the +// authoritative per-policy/group decision from management (the proxy per-provider +// backstop carries the union of both models), so this also confirms management +// receives the correct normalized model for path-routed providers. +func TestGuardrailPerGroupAllowlist_AllProviders(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Minute) + defer cancel() + + const ( + vertexProject = "e2e-project" + vertexRegion = "global" + ) + + priced := func(ids ...string) *[]api.AgentNetworkProviderModel { + out := make([]api.AgentNetworkProviderModel, 0, len(ids)) + for _, id := range ids { + out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001}) + } + return &out + } + + cases := []*pergroupCase{ + { + name: "openai", catalogID: "openai_api", wire: harness.WireChat, + models: priced("oai-model-a", "oai-model-b"), + selectedReq: "oai-model-a", selectedAllow: "oai-model-a", + otherReq: "oai-model-b", otherAllow: "oai-model-b", + }, + { + name: "anthropic", catalogID: "anthropic_api", wire: harness.WireMessages, + models: priced("ant-model-a", "ant-model-b"), + selectedReq: "ant-model-a", selectedAllow: "ant-model-a", + otherReq: "ant-model-b", otherAllow: "ant-model-b", + }, + { + // Vertex catalog ids travel bare in the rawPredict path. + name: "vertex", catalogID: "vertex_ai_api", wire: "vertex", + selectedReq: "claude-sonnet-4-5", selectedAllow: "claude-sonnet-4-5", + otherReq: "claude-opus-4-6", otherAllow: "claude-opus-4-6", + }, + { + // Bedrock request ids are region-prefixed/versioned; the parser + // normalizes them to the catalog key the allowlist holds. + name: "bedrock", catalogID: "bedrock_api", wire: "bedrock", + selectedReq: "us.anthropic.claude-sonnet-4-5-v1:0", selectedAllow: "anthropic.claude-sonnet-4-5", + otherReq: "us.anthropic.claude-opus-4-8-v1:0", otherAllow: "anthropic.claude-opus-4-8", + }, + } + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-pergroup-main"}) + require.NoError(t, err, "create main group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) }) + + grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-pergroup-other"}) + require.NoError(t, err, "create other group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-pergroup-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grpMain.Id}, + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + staticKey := "static-e2e-token" + enabled := true + + for i, c := range cases { + req := api.AgentNetworkProviderRequest{ + Name: "e2e-pergroup-" + c.name, + ProviderId: c.catalogID, + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: c.models, + } + if i == 0 { + req.BootstrapCluster = ptr(harness.AgentNetworkCluster) + } + prov, perr := srv.CreateProvider(ctx, req) + require.NoError(t, perr, "create provider %s", c.name) + c.providerID = prov.Id + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) }) + + gSel := mkAllowGuardrail(t, ctx, "e2e-pergroup-"+c.name+"-sel", c.selectedAllow) + gOth := mkAllowGuardrail(t, ctx, "e2e-pergroup-"+c.name+"-oth", c.otherAllow) + + polMain, merr := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-pergroup-" + c.name + "-main", + Enabled: &enabled, + SourceGroups: []string{grpMain.Id}, + DestinationProviderIds: []string{prov.Id}, + GuardrailIds: &[]string{gSel.Id}, + }) + require.NoError(t, merr, "create main policy %s", c.name) + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) }) + + polOther, oerr := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-pergroup-" + c.name + "-other", + Enabled: &enabled, + SourceGroups: []string{grpOther.Id}, + DestinationProviderIds: []string{prov.Id}, + GuardrailIds: &[]string{gOth.Id}, + }) + require.NoError(t, oerr, "create other policy %s", c.name) + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) }) + } + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-pergroup-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + send := func(c *pergroupCase, model string) (int, string) { + var code int + var body string + var cerr error + switch c.wire { + case "vertex": + code, body, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, vertexProject, vertexRegion, model, "Reply with exactly: pong", "") + case "bedrock": + code, body, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, model, "Reply with exactly: pong", "") + default: + code, body, cerr = cl.Chat(ctx, settings.Endpoint, proxyIP, c.wire, model, "Reply with exactly: pong", "") + } + require.NoError(t, cerr, "request must reach the proxy for %s", c.name) + return code, body + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + // grpMain's own model is served. Retry to absorb tunnel/DNS jitter on + // the first call over the freshly warmed tunnel. + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + code, body = send(c, c.selectedReq) + if code == 200 { + break + } + time.Sleep(5 * time.Second) + } + assert.Equal(t, 200, code, + "%s: grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", c.name, body, px.Logs(context.Background())) + + // grpOther's model must NOT leak to the grpMain client. + code, body = send(c, c.otherReq) + assert.Equal(t, 403, code, + "%s: another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", c.name, body, px.Logs(context.Background())) + assert.Contains(t, body, "llm_policy.model_blocked", + "%s: denial must be a model-allowlist decision, not routing; body: %s", c.name, body) + }) + } +} + +// TestGuardrailMultiGroupUser proves the per-policy/group decision for a caller +// that belongs to MULTIPLE groups at once. Two scenarios, one shared stack: +// +// - union across the user's groups: the client is in gUX and gUY, each with +// its own policy+guardrail on provider P1 (gUX->union-a, gUY->union-b). The +// client may use BOTH models (the union of its groups' allowlists) while a +// third, un-allowlisted model is denied. +// - an un-guardrailed group lifts the restriction: the client is in gMP and +// gMQ on provider P2, where gMP restricts to mix-a but gMQ's policy carries +// NO guardrail. Because one applicable policy is unrestricted, the client may +// use a model on no allowlist (mix-z) as well as mix-a. +func TestGuardrailMultiGroupUser(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + unionA = "mg-union-a" + unionB = "mg-union-b" + unionC = "mg-union-c" // allowlisted by neither group + mixA = "mg-mix-a" + mixZ = "mg-mix-z" // on no allowlist; reachable only via the un-guardrailed policy + ) + + priced := func(ids ...string) *[]api.AgentNetworkProviderModel { + out := make([]api.AgentNetworkProviderModel, 0, len(ids)) + for _, id := range ids { + out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001}) + } + return &out + } + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + mkGroup := func(name string) *api.Group { + g, gerr := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: name}) + require.NoError(t, gerr, "create group %s", name) + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), g.Id) }) + return g + } + gUX := mkGroup("e2e-mg-union-x") + gUY := mkGroup("e2e-mg-union-y") + gMP := mkGroup("e2e-mg-mix-p") + gMQ := mkGroup("e2e-mg-mix-q") + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-mg-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{gUX.Id, gUY.Id, gMP.Id, gMQ.Id}, // client in all four groups + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + staticKey := "static-e2e-token" + enabled := true + + // P1 — union scenario: two restricting policies, one per group. + p1, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "e2e-mg-union", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: priced(unionA, unionB, unionC), + BootstrapCluster: ptr(harness.AgentNetworkCluster), + }) + require.NoError(t, err, "create union provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p1.Id) }) + + polUX, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-mg-union-x", + Enabled: &enabled, + SourceGroups: []string{gUX.Id}, + DestinationProviderIds: []string{p1.Id}, + GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-union-x", unionA).Id}, + }) + require.NoError(t, err, "create union policy X") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polUX.Id) }) + + polUY, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-mg-union-y", + Enabled: &enabled, + SourceGroups: []string{gUY.Id}, + DestinationProviderIds: []string{p1.Id}, + GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-union-y", unionB).Id}, + }) + require.NoError(t, err, "create union policy Y") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polUY.Id) }) + + // P2 — mixed scenario: one restricting policy + one un-guardrailed policy. + p2, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: "e2e-mg-mix", + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &staticKey, + Enabled: ptr(true), + Models: priced(mixA, mixZ), + }) + require.NoError(t, err, "create mix provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p2.Id) }) + + polMP, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-mg-mix-p", + Enabled: &enabled, + SourceGroups: []string{gMP.Id}, + DestinationProviderIds: []string{p2.Id}, + GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-mix-p", mixA).Id}, + }) + require.NoError(t, err, "create mix policy P") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMP.Id) }) + + polMQ, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-mg-mix-q", + Enabled: &enabled, + SourceGroups: []string{gMQ.Id}, + DestinationProviderIds: []string{p2.Id}, // NO guardrail -> unrestricted + }) + require.NoError(t, err, "create mix policy Q") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMQ.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-mg-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + send := func(model string) (int, string) { + code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "") + require.NoError(t, cerr, "request must reach the proxy") + return code, body + } + sendUntil200 := func(model string) (int, string) { + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + code, body = send(model) + if code == 200 { + break + } + time.Sleep(5 * time.Second) + } + return code, body + } + + t.Run("union across the user's groups", func(t *testing.T) { + code, body := sendUntil200(unionA) + assert.Equal(t, 200, code, "model allowed by group X must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + + code, body = sendUntil200(unionB) + assert.Equal(t, 200, code, "model allowed by group Y must also be served (union across the user's groups); body: %s", body) + + code, body = send(unionC) + assert.Equal(t, 403, code, "a model on neither group's allowlist must be denied; body: %s", body) + assert.Contains(t, body, "llm_policy.model_blocked") + }) + + t.Run("an un-guardrailed group lifts the restriction", func(t *testing.T) { + code, body := sendUntil200(mixA) + assert.Equal(t, 200, code, "the restricted group's model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + + code, body = sendUntil200(mixZ) + assert.Equal(t, 200, code, + "a non-allowlisted model must be served because the user is also in a group whose policy has no guardrail; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background())) + }) +} + +// mkAllowGuardrail creates a guardrail whose model allowlist is enabled and holds +// exactly the given model, registering cleanup. +func mkAllowGuardrail(t *testing.T, ctx context.Context, name, model string) api.AgentNetworkGuardrail { + t.Helper() + var gr api.AgentNetworkGuardrailRequest + gr.Name = name + gr.Checks.ModelAllowlist.Enabled = true + gr.Checks.ModelAllowlist.Models = []string{model} + g, err := srv.CreateGuardrail(ctx, gr) + require.NoError(t, err, "create guardrail %s", name) + t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) }) + return g +} diff --git a/e2e/harness/proxy.go b/e2e/harness/proxy.go index 8db2c140f..85f3518d4 100644 --- a/e2e/harness/proxy.go +++ b/e2e/harness/proxy.go @@ -38,7 +38,11 @@ type Proxy struct { // network, registered via the given account proxy token and serving the // AgentNetworkCluster over a self-signed wildcard cert. It does not wait for // peer connectivity — callers poll management for the proxy peer. -func StartProxy(ctx context.Context, c *Combined, proxyToken string) (*Proxy, error) { +// StartProxy launches the reverse-proxy container. Optional envOverrides are +// merged into the container environment after the defaults, so callers can set +// or override any NB_PROXY_* var (e.g. NB_PROXY_TUNNEL_CACHE_TTL for tests that +// need a short authorization-cache window). +func StartProxy(ctx context.Context, c *Combined, proxyToken string, envOverrides ...map[string]string) (*Proxy, error) { root, err := repoRoot() if err != nil { return nil, err @@ -93,6 +97,12 @@ func StartProxy(ctx context.Context, c *Combined, proxyToken string) (*Proxy, er WaitingFor: wait.ForLog("Initial mapping sync complete").WithStartupTimeout(90 * time.Second), } + for _, ov := range envOverrides { + for k, v := range ov { + req.Env[k] = v + } + } + ctr, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ ContainerRequest: req, Started: true, diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index d88e0d77c..77c77ce44 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -86,12 +86,17 @@ type Manager interface { // PolicySelectionInput is the per-request selection envelope. The // proxy populates it from CapturedData (account, user, groups) plus -// the provider llm_router resolved. +// the provider llm_router resolved and the model it extracted. type PolicySelectionInput struct { AccountID string UserID string GroupIDs []string ProviderID string + // Model is the already-normalised upstream model id the proxy extracted + // (parser strips Bedrock region/version, Vertex @version), so a + // case-insensitive compare suffices. Empty = undetermined → not permitted + // (fail closed). + Model string } // PolicySelectionResult names the policy that "pays" for this request diff --git a/management/internals/modules/agentnetwork/policyselect.go b/management/internals/modules/agentnetwork/policyselect.go index 9203a1910..9bb893b36 100644 --- a/management/internals/modules/agentnetwork/policyselect.go +++ b/management/internals/modules/agentnetwork/policyselect.go @@ -5,6 +5,7 @@ import ( "fmt" "math" "sort" + "strings" "time" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" @@ -35,6 +36,10 @@ const ( denyCodeAccountTokenCapExceeded = "llm_account.token_cap_exceeded" //nolint:gosec // account deny code label, not a credential denyCodeAccountBudgetCapExceeded = "llm_account.budget_cap_exceeded" + // denyCodeModelBlocked is returned when policies govern the request's + // (provider, caller-groups) but none permits the model. Matches the proxy + // guardrail's code so both layers surface the same label. + denyCodeModelBlocked = "llm_policy.model_blocked" ) // consumptionCache holds the consumption counters prefetched for one @@ -159,6 +164,25 @@ func (m *managerImpl) SelectPolicyForRequest(ctx context.Context, in PolicySelec } candidates := filterApplicablePolicies(policies, in) + // Model-allowlist gate scoped to the matched policies: keep candidates whose + // guardrails permit the model (none enabled = unrestricted), deny when + // policies apply but none permits it. Skip the load when none has a guardrail. + if len(candidates) > 0 && anyPolicyHasGuardrails(candidates) { + guardrailsByID, gErr := m.loadGuardrailsByID(ctx, in.AccountID) + if gErr != nil { + return nil, gErr + } + permitted := filterModelPermittedPolicies(candidates, guardrailsByID, in.Model) + if len(permitted) == 0 { + return &PolicySelectionResult{ + Allow: false, + DenyCode: denyCodeModelBlocked, + DenyReason: modelBlockedReason(in.Model), + }, nil + } + candidates = permitted + } + // Prefetch every consumption counter the ceiling + candidate policies will // read, in a single store round-trip, then score against the cache. cache, err := m.prefetchConsumption(ctx, in, rules, candidates, now) @@ -250,6 +274,90 @@ func filterApplicablePolicies(policies []*types.Policy, in PolicySelectionInput) return out } +// anyPolicyHasGuardrails reports whether any policy references at least one +// guardrail, so the selector can skip loading guardrails when none do. +func anyPolicyHasGuardrails(policies []*types.Policy) bool { + for _, p := range policies { + if p != nil && len(p.GuardrailIDs) > 0 { + return true + } + } + return false +} + +// loadGuardrailsByID loads the account's guardrails indexed by ID. Used by the +// model-allowlist gate to resolve each candidate policy's attached guardrails. +func (m *managerImpl) loadGuardrailsByID(ctx context.Context, accountID string) (map[string]*types.Guardrail, error) { + guardrails, err := m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID) + if err != nil { + return nil, fmt.Errorf("list account guardrails: %w", err) + } + byID := make(map[string]*types.Guardrail, len(guardrails)) + for _, g := range guardrails { + if g != nil { + byID[g.ID] = g + } + } + return byID, nil +} + +// filterModelPermittedPolicies returns the subset of policies whose guardrails +// permit the model. Order is preserved so downstream scoring is unaffected. +func filterModelPermittedPolicies(policies []*types.Policy, byID map[string]*types.Guardrail, model string) []*types.Policy { + out := make([]*types.Policy, 0, len(policies)) + for _, p := range policies { + if policyPermitsModel(p, byID, model) { + out = append(out, p) + } + } + return out +} + +// policyPermitsModel reports whether a policy permits the model. No +// allowlist-enabled guardrail = unrestricted (permits any, incl. empty); +// otherwise the model must be in the union of its allowlists, so an +// empty/undetermined model fails closed. +func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model string) bool { + if p == nil { + return false + } + wanted := normaliseModelID(model) + restricted := false + for _, gID := range p.GuardrailIDs { + g, ok := byID[gID] + if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled { + continue + } + restricted = true + if wanted == "" { + continue + } + for _, allowed := range g.Checks.ModelAllowlist.Models { + if normaliseModelID(allowed) == wanted { + return true + } + } + } + return !restricted +} + +// normaliseModelID lowercases and trims a model identifier so the allowlist +// compare is case-insensitive and trim-tolerant. Mirrors the proxy guardrail's +// normaliseModel so both layers agree on what "same model" means. +func normaliseModelID(model string) string { + return strings.ToLower(strings.TrimSpace(model)) +} + +// modelBlockedReason builds the human-readable deny reason for a model-allowlist +// rejection. The model is quoted when known; an undetermined model is reported +// as such so the access log distinguishes "wrong model" from "no model". +func modelBlockedReason(model string) string { + if normaliseModelID(model) == "" { + return "request model could not be determined for the policy allowlist" + } + return fmt.Sprintf("model %q is not permitted by any applicable policy allowlist", model) +} + // candidate is the per-policy intermediate the selector ranks. A // policy that's been exhausted on any enabled cap never makes it // into this slice; the selector's deny envelope carries the latest diff --git a/management/internals/modules/agentnetwork/policyselect_model_test.go b/management/internals/modules/agentnetwork/policyselect_model_test.go new file mode 100644 index 000000000..c122cc36c --- /dev/null +++ b/management/internals/modules/agentnetwork/policyselect_model_test.go @@ -0,0 +1,329 @@ +package agentnetwork + +import ( + "context" + "errors" + "testing" + "time" + + "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/store" +) + +// guardedPolicy builds an enabled, uncapped policy that authorises sourceGroups +// to reach providerID under the given guardrails. Uncapped keeps the selector's +// headroom scoring trivial so these tests isolate the model-allowlist gate. +func guardedPolicy(id, account string, sourceGroups []string, providerID string, guardrailIDs ...string) *types.Policy { + return &types.Policy{ + ID: id, + AccountID: account, + Enabled: true, + SourceGroups: sourceGroups, + DestinationProviderIDs: []string{providerID}, + GuardrailIDs: guardrailIDs, + CreatedAt: time.Now().UTC(), + } +} + +// allowlistGuardrail builds a guardrail whose model allowlist is enabled and +// carries the given models. +func allowlistGuardrail(id, account string, models ...string) *types.Guardrail { + return &types.Guardrail{ + ID: id, + AccountID: account, + Checks: types.GuardrailChecks{ + ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true, Models: models}, + }, + } +} + +func expectPolicies(mockStore *store.MockStore, account string, policies ...*types.Policy) { + mockStore.EXPECT(). + GetAccountAgentNetworkPolicies(gomock.Any(), gomock.Any(), account). + Return(policies, nil) +} + +func expectGuardrails(mockStore *store.MockStore, account string, guardrails ...*types.Guardrail) { + mockStore.EXPECT(). + GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), account). + Return(guardrails, nil) +} + +// TestSelectPolicy_ModelBlockedByAllowlist proves the authoritative allowlist +// decision: a policy authorises the (provider, group) but restricts the model, +// and the requested model isn't on the list, so the request is denied. +func TestSelectPolicy_ModelBlockedByAllowlist(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o")) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + UserID: "user-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + assert.False(t, res.Allow, "a model outside the only applicable policy's allowlist must be denied") + assert.Equal(t, denyCodeModelBlocked, res.DenyCode, "deny code must be model_blocked") + assert.NotEmpty(t, res.DenyReason, "deny reason must be populated") +} + +// TestSelectPolicy_ModelAllowedByAllowlist is the allow counterpart: the model +// is on the applicable policy's allowlist, so selection proceeds normally. +func TestSelectPolicy_ModelAllowedByAllowlist(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o", "claude-opus-4")) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + UserID: "user-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "a model on the applicable policy's allowlist must be allowed") + assert.Equal(t, "pol-A", res.SelectedPolicyID) +} + +// TestSelectPolicy_CaseInsensitiveModelMatch proves the compare tolerates case +// and surrounding whitespace, matching the proxy guardrail's normalisation. +func TestSelectPolicy_CaseInsensitiveModelMatch(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", " GPT-4o ")) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "gpt-4o", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "case/whitespace variants must match the allowlist entry") +} + +// TestSelectPolicy_UnguardedPolicyIsUnrestricted is the false-deny fix: when two +// policies authorise the same (provider, group) and one has no guardrail, that +// policy makes the request unrestricted — not caught by the other's allowlist. +func TestSelectPolicy_UnguardedPolicyIsUnrestricted(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + restricted := guardedPolicy("pol-restricted", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + open := guardedPolicy("pol-open", "acc-1", []string{"grp-eng"}, "prov-1") // no guardrail + expectPolicies(mockStore, "acc-1", restricted, open) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o")) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "an un-guardrailed policy for the same (provider, group) must leave the request unrestricted") + assert.Equal(t, "pol-open", res.SelectedPolicyID, "the unrestricted policy must be the one that pays") +} + +// TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups is the false-allow fix: a +// model allowlisted only for grp-b must not be usable by a grp-a caller. The +// selector considers only policies applicable to the caller's groups. +func TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + polA := guardedPolicy("pol-a", "acc-1", []string{"grp-a"}, "prov-1", "g-a") + polB := guardedPolicy("pol-b", "acc-1", []string{"grp-b"}, "prov-1", "g-b") + expectPolicies(mockStore, "acc-1", polA, polB) + expectGuardrails(mockStore, "acc-1", + allowlistGuardrail("g-a", "acc-1", "gpt-4o"), + allowlistGuardrail("g-b", "acc-1", "claude-opus-4"), + ) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-a"}, + ProviderID: "prov-1", + Model: "claude-opus-4", // only allowed for grp-b + }) + require.NoError(t, err) + assert.False(t, res.Allow, "grp-b's allowlisted model must not leak to a grp-a caller") + assert.Equal(t, denyCodeModelBlocked, res.DenyCode) +} + +// TestSelectPolicy_UndeterminedModelFailsClosed proves the fail-closed contract +// mirrors the proxy: with a restricted applicable policy and an empty model +// (e.g. a path-routed shape the parser couldn't map), the request is denied. +func TestSelectPolicy_UndeterminedModelFailsClosed(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o")) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "", // undetermined + }) + require.NoError(t, err) + assert.False(t, res.Allow, "an undetermined model must fail closed against a restricted policy") + assert.Equal(t, denyCodeModelBlocked, res.DenyCode) +} + +// TestSelectPolicy_DisabledAllowlistDoesNotRestrict proves a guardrail whose +// model allowlist is disabled imposes no model restriction, even though the +// policy references it. +func TestSelectPolicy_DisabledAllowlistDoesNotRestrict(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + disabled := &types.Guardrail{ + ID: "g-1", + AccountID: "acc-1", + Checks: types.GuardrailChecks{ + ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}, + }, + } + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", disabled) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "anything-goes", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "a disabled allowlist must not restrict the model") + assert.Equal(t, "pol-A", res.SelectedPolicyID) +} + +// TestSelectPolicy_UnionAcrossPolicyGuardrails proves a policy with multiple +// allowlist guardrails permits the union of their models (not just the first). +func TestSelectPolicy_UnionAcrossPolicyGuardrails(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1", "g-2") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1", + allowlistGuardrail("g-1", "acc-1", "gpt-4o"), + allowlistGuardrail("g-2", "acc-1", "claude-opus-4"), + ) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", // only in the second guardrail's list + }) + require.NoError(t, err) + assert.True(t, res.Allow, "a model in any of the policy's allowlist guardrails must be permitted") + assert.Equal(t, "pol-A", res.SelectedPolicyID) +} + +// TestSelectPolicy_GuardrailLookupErrorPropagates proves a store failure while +// resolving the candidate policies' guardrails surfaces as an error, not a +// silent allow/deny. +func TestSelectPolicy_GuardrailLookupErrorPropagates(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1") + expectPolicies(mockStore, "acc-1", policy) + mockStore.EXPECT(). + GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), "acc-1"). + Return(nil, errors.New("store unavailable")) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "gpt-4o", + }) + require.Error(t, err, "a guardrail-lookup failure must surface as an error") + assert.Nil(t, res) +} + +// TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted proves a +// policy referencing a guardrail ID absent from the account's set (a stale +// reference) imposes no model restriction — same as no guardrail. +func TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-missing") + expectPolicies(mockStore, "acc-1", policy) + expectGuardrails(mockStore, "acc-1") + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "anything-goes", + }) + require.NoError(t, err) + assert.True(t, res.Allow, "an orphaned guardrail reference must not restrict the model") + assert.Equal(t, "pol-A", res.SelectedPolicyID) +} + +// TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter proves the model +// gate narrows candidates before cap scoring: the permitting policy is selected +// even though the blocked one has a larger, more attractive cap. +func TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter(t *testing.T) { + ctrl := gomock.NewController(t) + mgr, mockStore := newSelectorMgr(t, ctrl) + + polBig := guardedPolicy("pol-big", "acc-1", []string{"grp-eng"}, "prov-1", "g-restrict") + polBig.Limits = types.PolicyLimits{ + TokenLimit: types.PolicyTokenLimit{Enabled: true, GroupCap: 1_000_000, WindowSeconds: 3600}, + } + polSmall := guardedPolicy("pol-small", "acc-1", []string{"grp-eng"}, "prov-1", "g-permit") + polSmall.Limits = types.PolicyLimits{ + TokenLimit: types.PolicyTokenLimit{Enabled: true, GroupCap: 100, WindowSeconds: 3600}, + } + expectPolicies(mockStore, "acc-1", polBig, polSmall) + expectGuardrails(mockStore, "acc-1", + allowlistGuardrail("g-restrict", "acc-1", "gpt-4o"), + allowlistGuardrail("g-permit", "acc-1", "claude-opus-4"), + ) + expectConsumptionBatch(mockStore, nil) + + res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{ + AccountID: "acc-1", + GroupIDs: []string{"grp-eng"}, + ProviderID: "prov-1", + Model: "claude-opus-4", // only pol-small's guardrail permits this + }) + require.NoError(t, err) + assert.True(t, res.Allow) + assert.Equal(t, "pol-small", res.SelectedPolicyID, + "the model filter must exclude pol-big before cap scoring") +} diff --git a/management/internals/modules/agentnetwork/synthesizer.go b/management/internals/modules/agentnetwork/synthesizer.go index 95fe91773..169bdd4fd 100644 --- a/management/internals/modules/agentnetwork/synthesizer.go +++ b/management/internals/modules/agentnetwork/synthesizer.go @@ -235,7 +235,12 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([ mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID) applyAccountCollectionControls(&mergedGuardrails, settings) - guardrailJSON, err := marshalGuardrailConfig(mergedGuardrails) + // The proxy guardrail is a per-provider fail-closed backstop; the + // authoritative per-policy/group decision is management's + // SelectPolicyForRequest. A provider lands in this map only when every + // authorising policy restricts models. + providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID) + guardrailJSON, err := marshalGuardrailConfig(providerAllowlists, mergedGuardrails.PromptCapture) if err != nil { return nil, err } @@ -780,10 +785,12 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt // guardrailConfig is the JSON shape the proxy-side llm_guardrail // middleware expects. Mirrors the proxy registration documented in -// the management→proxy contract. +// the management→proxy contract. provider_allowlists is keyed by the +// resolved provider id llm_router stamps; a provider absent from the map is +// unrestricted at the proxy layer. type guardrailConfig struct { - ModelAllowlist []string `json:"model_allowlist,omitempty"` - PromptCapture guardrailPromptCapture `json:"prompt_capture"` + ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"` + PromptCapture guardrailPromptCapture `json:"prompt_capture"` } type guardrailPromptCapture struct { @@ -828,13 +835,10 @@ func applyAccountCollectionControls(merged *MergedGuardrails, settings *types.Se merged.PromptCapture.RedactPii = settings.RedactPii || merged.PromptCapture.RedactPii } -func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) { +func marshalGuardrailConfig(providerAllowlists map[string][]string, capture MergedPromptCapture) ([]byte, error) { cfg := guardrailConfig{ - ModelAllowlist: merged.ModelAllowlist, - PromptCapture: guardrailPromptCapture{ - Enabled: merged.PromptCapture.Enabled, - RedactPii: merged.PromptCapture.RedactPii, - }, + ProviderAllowlists: providerAllowlists, + PromptCapture: guardrailPromptCapture(capture), } out, err := json.Marshal(cfg) if err != nil { @@ -843,6 +847,74 @@ func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) { return out, nil } +// buildProviderAllowlists returns the proxy's per-provider backstop: a provider +// is included only when every authorising policy restricts models (their union); +// if any leaves it unrestricted it is omitted, so management decides per group. +func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Guardrail) map[string][]string { + type providerAcc struct { + models map[string]struct{} + anyUnrestricted bool + } + accs := make(map[string]*providerAcc) + for _, p := range policies { + if p == nil { + continue + } + restricted, models := policyModelAllowlist(p, byID) + for _, providerID := range p.DestinationProviderIDs { + if providerID == "" { + continue + } + acc, ok := accs[providerID] + if !ok { + acc = &providerAcc{models: make(map[string]struct{})} + accs[providerID] = acc + } + if !restricted { + acc.anyUnrestricted = true + continue + } + for _, m := range models { + acc.models[m] = struct{}{} + } + } + } + out := make(map[string][]string, len(accs)) + for providerID, acc := range accs { + if acc.anyUnrestricted { + continue + } + models := make([]string, 0, len(acc.models)) + for m := range acc.models { + models = append(models, m) + } + sort.Strings(models) + out[providerID] = models + } + return out +} + +// policyModelAllowlist reports whether a policy restricts models (has an +// allowlist-enabled guardrail) and the union of allowed models. Models are +// verbatim; the proxy factory lowercases/trims them at decode time. +func policyModelAllowlist(p *types.Policy, byID map[string]*types.Guardrail) (bool, []string) { + restricted := false + var models []string + for _, gID := range p.GuardrailIDs { + g, ok := byID[gID] + if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled { + continue + } + restricted = true + for _, m := range g.Checks.ModelAllowlist.Models { + if m != "" { + models = append(models, m) + } + } + } + return restricted, models +} + // buildAccountService composes the per-account gateway Service. The // target carries the noop placeholder URL — the router middleware // rewrites every request to the matched provider's upstream before the @@ -986,38 +1058,11 @@ func unionSourceGroups(policies []*types.Policy) []string { return out } -// MergedGuardrails is the JSON shape passed to the proxy via the -// guardrail middleware's config_json. Mirrors the proxy-side -// expectations and is intentionally distinct from -// types.GuardrailChecks so we can evolve either side independently. +// MergedGuardrails is the synthesiser's fold target. Only prompt capture is +// merged here — the model allowlist is emitted per-provider, and +// token/budget/retention moved onto Policy.Limits and account Settings. type MergedGuardrails struct { - ModelAllowlist []string `json:"model_allowlist,omitempty"` - TokenLimits MergedTokenLimits `json:"token_limits"` - Budget MergedBudget `json:"budget"` - PromptCapture MergedPromptCapture `json:"prompt_capture"` - Retention MergedRetention `json:"retention"` -} - -type MergedTokenLimits struct { - Hourly *MergedTokenWindow `json:"hourly,omitempty"` - Daily *MergedTokenWindow `json:"daily,omitempty"` - Monthly *MergedTokenWindow `json:"monthly,omitempty"` -} - -type MergedTokenWindow struct { - MaxInputTokens int `json:"max_input_tokens,omitempty"` - MaxOutputTokens int `json:"max_output_tokens,omitempty"` -} - -type MergedBudget struct { - Hourly *MergedBudgetWindow `json:"hourly,omitempty"` - Daily *MergedBudgetWindow `json:"daily,omitempty"` - Monthly *MergedBudgetWindow `json:"monthly,omitempty"` -} - -type MergedBudgetWindow struct { - SoftCapUSD float64 `json:"soft_cap_usd,omitempty"` - HardCapUSD float64 `json:"hard_cap_usd,omitempty"` + PromptCapture MergedPromptCapture } type MergedPromptCapture struct { @@ -1025,64 +1070,31 @@ type MergedPromptCapture struct { RedactPii bool `json:"redact_pii"` } -type MergedRetention struct { - Enabled bool `json:"enabled"` - Days int `json:"days"` -} - -// mergeGuardrails computes the effective guardrail spec applied at the -// proxy, given the referencing policies and the account's guardrail -// catalogue. Policy enabled-ness is the caller's responsibility — only -// enabled policies should be passed in. +// mergeGuardrails folds the referencing policies' guardrails into the +// prompt-capture decision only. The model allowlist is enforced per-policy/group +// in management and shipped per-provider; token/budget/retention live off +// guardrails now. // -// Merge rules: -// - Model allowlist: union of allowlists across policies that enable it. -// - Token / Budget: most-restrictive (min of non-zero caps) per window. -// - Prompt capture: enabled if any policy enables it; redact_pii sticks -// if any enabling policy turns it on. -// - Retention: enabled if any enables it; smallest non-zero days wins. +// Merge rule — prompt capture: enabled if any policy enables it; redact_pii +// sticks if any enabling policy turns it on. func mergeGuardrails(policies []*types.Policy, byID map[string]*types.Guardrail) MergedGuardrails { merged := MergedGuardrails{} - allowlist := make(map[string]struct{}) - allowlistEnabled := false - for _, policy := range policies { for _, gID := range policy.GuardrailIDs { g, ok := byID[gID] if !ok || g == nil { continue } - mergeGuardrail(g, &merged, allowlist, &allowlistEnabled) + mergeGuardrail(g, &merged) } } - - if allowlistEnabled { - merged.ModelAllowlist = make([]string, 0, len(allowlist)) - for m := range allowlist { - merged.ModelAllowlist = append(merged.ModelAllowlist, m) - } - sort.Strings(merged.ModelAllowlist) - } return merged } -// mergeGuardrail folds a single guardrail's enabled checks into the -// running merge: model-allowlist models join the shared set (and flip -// allowlistEnabled), and prompt-capture / redact-pii stick once any -// enabling guardrail turns them on. -// -// TokenLimits, Budget, and Retention have moved off guardrails — token -// and budget caps now live on the Policy itself (Policy.Limits) and -// retention moves to account-level Settings — so they are not merged here. -func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails, allowlist map[string]struct{}, allowlistEnabled *bool) { - if g.Checks.ModelAllowlist.Enabled { - *allowlistEnabled = true - for _, m := range g.Checks.ModelAllowlist.Models { - if m != "" { - allowlist[m] = struct{}{} - } - } - } +// mergeGuardrail folds a single guardrail's prompt-capture settings into the +// running merge: enabled / redact-pii stick once any enabling guardrail turns +// them on. +func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) { if g.Checks.PromptCapture.Enabled { merged.PromptCapture.Enabled = true if g.Checks.PromptCapture.RedactPii { diff --git a/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go b/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go index 8ed4910da..8b13f78ac 100644 --- a/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_guardrail_realstore_test.go @@ -81,8 +81,8 @@ func TestSynthesizeServices_RealStore_PromptCaptureAccountIsSoleControl(t *testi require.Len(t, services, 1) cfg := decodeServiceGuardrailConfig(t, services[0]) - assert.Equal(t, []string{"gpt-5.4"}, cfg.ModelAllowlist, - "model allowlist is a pure policy guardrail and must always reach the config") + assert.Equal(t, map[string][]string{"prov-1": {"gpt-5.4"}}, cfg.ProviderAllowlists, + "model allowlist is a pure policy guardrail and must reach the per-provider config") assert.False(t, cfg.PromptCapture.Enabled, "prompt capture must be off when the account toggle is off, even with a capture-enabled guardrail") } @@ -172,7 +172,7 @@ func TestSynthesizeServices_RealStore_NoGuardrail_CaptureOff(t *testing.T) { require.Len(t, services, 1, "exactly one synth service expected") cfg := decodeServiceGuardrailConfig(t, services[0]) - assert.Empty(t, cfg.ModelAllowlist, "no guardrail → no allowlist") + assert.Empty(t, cfg.ProviderAllowlists, "no guardrail → provider unrestricted (absent from map)") assert.False(t, cfg.PromptCapture.Enabled, "no guardrail → prompt capture off by default") assert.False(t, cfg.PromptCapture.RedactPii, "no guardrail → redact off by default") } diff --git a/management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go b/management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go new file mode 100644 index 000000000..2cfc0db8c --- /dev/null +++ b/management/internals/modules/agentnetwork/synthesizer_provider_allowlist_test.go @@ -0,0 +1,95 @@ +package agentnetwork + +import ( + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" +) + +// policyForProviders builds an enabled policy authorising the given providers +// under the given guardrails (both optional). Groups are irrelevant to +// buildProviderAllowlists, which keys purely on destination provider. +func policyForProviders(id string, guardrailIDs []string, providerIDs ...string) *types.Policy { + return &types.Policy{ + ID: id, + Enabled: true, + DestinationProviderIDs: providerIDs, + GuardrailIDs: guardrailIDs, + } +} + +func TestBuildProviderAllowlists(t *testing.T) { + byID := map[string]*types.Guardrail{ + "g-4o": allowlistGuardrail("g-4o", "acc-1", "gpt-4o"), + "g-opus": allowlistGuardrail("g-opus", "acc-1", "claude-opus-4"), + "g-disabled": {ID: "g-disabled", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}}}, + } + + t.Run("all authorising policies restrict yields per-provider union", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x"), + policyForProviders("p2", []string{"g-opus"}, "prov-x"), + } + got := buildProviderAllowlists(policies, byID) + assert.Equal(t, map[string][]string{"prov-x": {"claude-opus-4", "gpt-4o"}}, got, + "a provider every policy restricts carries the sorted union of their models") + }) + + t.Run("any un-guardrailed policy leaves the provider unrestricted (omitted)", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x"), + policyForProviders("p2", nil, "prov-x"), // no guardrail + } + got := buildProviderAllowlists(policies, byID) + assert.NotContains(t, got, "prov-x", + "a provider reachable by an un-guardrailed policy must be omitted so the proxy treats it as unrestricted") + }) + + t.Run("a disabled allowlist counts as unrestricted", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-disabled"}, "prov-x"), + } + got := buildProviderAllowlists(policies, byID) + assert.NotContains(t, got, "prov-x", + "a policy whose only guardrail has a disabled allowlist is unrestricted") + }) + + t.Run("providers are isolated from one another", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x"), + policyForProviders("p2", []string{"g-opus"}, "prov-y"), + } + got := buildProviderAllowlists(policies, byID) + assert.Equal(t, []string{"gpt-4o"}, got["prov-x"], "prov-x keeps only its own model") + assert.Equal(t, []string{"claude-opus-4"}, got["prov-y"], "prov-y keeps only its own model") + }) + + t.Run("one policy authorising two providers restricts both", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o"}, "prov-x", "prov-y"), + } + got := buildProviderAllowlists(policies, byID) + assert.Equal(t, []string{"gpt-4o"}, got["prov-x"]) + assert.Equal(t, []string{"gpt-4o"}, got["prov-y"]) + }) + + t.Run("union across a single policy's guardrails", func(t *testing.T) { + policies := []*types.Policy{ + policyForProviders("p1", []string{"g-4o", "g-opus"}, "prov-x"), + } + got := buildProviderAllowlists(policies, byID) + assert.ElementsMatch(t, []string{"claude-opus-4", "gpt-4o"}, got["prov-x"], + "a policy's own multiple allowlist guardrails union together") + }) + + t.Run("an enabled allowlist with no models denies everything", func(t *testing.T) { + empty := map[string]*types.Guardrail{"g-empty": allowlistGuardrail("g-empty", "acc-1")} + got := buildProviderAllowlists([]*types.Policy{ + policyForProviders("p1", []string{"g-empty"}, "prov-x"), + }, empty) + assert.Equal(t, map[string][]string{"prov-x": {}}, got, + "an enabled-but-empty allowlist is restricted with an empty set, not unrestricted") + }) +} diff --git a/management/internals/modules/agentnetwork/synthesizer_test.go b/management/internals/modules/agentnetwork/synthesizer_test.go index 206d1d12a..8a18a9b59 100644 --- a/management/internals/modules/agentnetwork/synthesizer_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_test.go @@ -1031,8 +1031,12 @@ func TestSynthesizeServices_GuardrailMerge_AllowlistUnion_LimitsRestrictive(t *t var cfg guardrailConfig require.NoError(t, json.Unmarshal(guardrailJSON, &cfg), "guardrail config must unmarshal cleanly") - assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ModelAllowlist, - "model allowlist union must keep both models") + // Both policies restrict the same provider, so the per-provider backstop + // carries the union of their models — a coarse gate that management's + // per-policy/group check narrows; it only blocks models outside the union + // when management is down. + assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ProviderAllowlists["prov-1"], + "per-provider allowlist union must keep both models") } func TestSynthesizeServices_BackfillsMissingSessionKeys(t *testing.T) { diff --git a/management/internals/shared/grpc/proxy.go b/management/internals/shared/grpc/proxy.go index b289f8c71..8f24de116 100644 --- a/management/internals/shared/grpc/proxy.go +++ b/management/internals/shared/grpc/proxy.go @@ -285,6 +285,7 @@ func (s *ProxyServiceServer) CheckLLMPolicyLimits(ctx context.Context, req *prot UserID: req.GetUserId(), GroupIDs: req.GetGroupIds(), ProviderID: req.GetProviderId(), + Model: req.GetModel(), }) if err != nil { log.WithContext(ctx).Errorf("select policy for request: %v", err) diff --git a/management/internals/shared/grpc/proxy_llm_policy_limits_test.go b/management/internals/shared/grpc/proxy_llm_policy_limits_test.go new file mode 100644 index 000000000..c293c39ea --- /dev/null +++ b/management/internals/shared/grpc/proxy_llm_policy_limits_test.go @@ -0,0 +1,138 @@ +package grpc + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" + "github.com/netbirdio/netbird/shared/management/proto" +) + +// fakeAgentNetworkLimits records the PolicySelectionInput it was invoked with +// and returns a pre-programmed result, so tests can assert what the handler +// forwards to the selector. +type fakeAgentNetworkLimits struct { + gotInput agentnetwork.PolicySelectionInput + result *agentnetwork.PolicySelectionResult + err error +} + +func (f *fakeAgentNetworkLimits) SelectPolicyForRequest(_ context.Context, in agentnetwork.PolicySelectionInput) (*agentnetwork.PolicySelectionResult, error) { + f.gotInput = in + if f.err != nil { + return nil, f.err + } + return f.result, nil +} + +func (f *fakeAgentNetworkLimits) RecordUsage(_ context.Context, _ agentnetwork.RecordUsageInput) error { + return nil +} + +// TestCheckLLMPolicyLimits_ForwardsModelToSelector proves the wiring added here: +// the model the proxy extracted must reach the selector's Model unchanged, +// alongside the account/user/group/provider fields. +func TestCheckLLMPolicyLimits_ForwardsModelToSelector(t *testing.T) { + fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{Allow: true, SelectedPolicyID: "pol-1"}} + s := &ProxyServiceServer{} + s.SetAgentNetworkLimitsService(fake) + + req := &proto.CheckLLMPolicyLimitsRequest{ + AccountId: "acc-1", + UserId: "user-1", + GroupIds: []string{"grp-a", "grp-b"}, + ProviderId: "prov-1", + Model: "claude-opus-4", + } + + resp, err := s.CheckLLMPolicyLimits(context.Background(), req) + require.NoError(t, err) + require.NotNil(t, resp) + + assert.Equal(t, "acc-1", fake.gotInput.AccountID) + assert.Equal(t, "user-1", fake.gotInput.UserID) + assert.Equal(t, []string{"grp-a", "grp-b"}, fake.gotInput.GroupIDs) + assert.Equal(t, "prov-1", fake.gotInput.ProviderID) + assert.Equal(t, "claude-opus-4", fake.gotInput.Model, + "the request's model must be forwarded to the selector") +} + +// TestCheckLLMPolicyLimits_EmptyModelForwardedAsEmpty proves an undetermined +// model (empty string) is forwarded as-is; the selector decides how to treat it. +func TestCheckLLMPolicyLimits_EmptyModelForwardedAsEmpty(t *testing.T) { + fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{Allow: true}} + s := &ProxyServiceServer{} + s.SetAgentNetworkLimitsService(fake) + + req := &proto.CheckLLMPolicyLimitsRequest{ + AccountId: "acc-1", + ProviderId: "prov-1", + } + + _, err := s.CheckLLMPolicyLimits(context.Background(), req) + require.NoError(t, err) + assert.Equal(t, "", fake.gotInput.Model, "an absent model must be forwarded as empty, not fabricated") +} + +// TestCheckLLMPolicyLimits_DenyResponseCarriesModelBlockedCode proves the deny +// envelope surfaces the model-allowlist deny code + reason through the response. +func TestCheckLLMPolicyLimits_DenyResponseCarriesModelBlockedCode(t *testing.T) { + fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{ + Allow: false, + DenyCode: "llm_policy.model_blocked", + DenyReason: `model "claude-opus-4" is not permitted by any applicable policy allowlist`, + }} + s := &ProxyServiceServer{} + s.SetAgentNetworkLimitsService(fake) + + resp, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{ + AccountId: "acc-1", + ProviderId: "prov-1", + Model: "claude-opus-4", + }) + require.NoError(t, err) + require.NotNil(t, resp) + assert.Equal(t, "deny", resp.Decision) + assert.Equal(t, "llm_policy.model_blocked", resp.DenyCode) + assert.NotEmpty(t, resp.DenyReason) + assert.Empty(t, resp.SelectedPolicyId, "a denied request must carry no selected policy") +} + +// TestCheckLLMPolicyLimits_SelectorErrorSurfacesAsInternal proves a selector +// failure surfaces as an Internal gRPC error rather than a silent allow. +func TestCheckLLMPolicyLimits_SelectorErrorSurfacesAsInternal(t *testing.T) { + fake := &fakeAgentNetworkLimits{err: errors.New("boom")} + s := &ProxyServiceServer{} + s.SetAgentNetworkLimitsService(fake) + + _, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{ + AccountId: "acc-1", + ProviderId: "prov-1", + Model: "gpt-4o", + }) + require.Error(t, err) + st, ok := grpcstatus.FromError(err) + require.True(t, ok) + assert.Equal(t, codes.Internal, st.Code(), "selector errors must never fail open on the hot path") +} + +// TestCheckLLMPolicyLimits_UnconfiguredServiceReturnsUnimplemented locks the +// fallback: with no limits service wired the RPC returns Unimplemented. +func TestCheckLLMPolicyLimits_UnconfiguredServiceReturnsUnimplemented(t *testing.T) { + s := &ProxyServiceServer{} + + _, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{ + AccountId: "acc-1", + ProviderId: "prov-1", + }) + require.Error(t, err) + st, ok := grpcstatus.FromError(err) + require.True(t, ok) + assert.Equal(t, codes.Unimplemented, st.Code()) +} diff --git a/proxy/internal/auth/tunnel_cache.go b/proxy/internal/auth/tunnel_cache.go index 10b671d82..185c53c62 100644 --- a/proxy/internal/auth/tunnel_cache.go +++ b/proxy/internal/auth/tunnel_cache.go @@ -3,20 +3,30 @@ package auth import ( "context" "net/netip" + "os" + "strings" "sync" "time" + log "github.com/sirupsen/logrus" "golang.org/x/sync/singleflight" "github.com/netbirdio/netbird/proxy/internal/types" "github.com/netbirdio/netbird/shared/management/proto" ) -// tunnelCacheTTL caps how long a positive ValidateTunnelPeer result is -// reused before re-fetching from management. 5 minutes balances freshness -// against management load on busy mesh networks. +// tunnelCacheTTL is the default cap on how long a positive ValidateTunnelPeer +// result is reused before re-fetching from management. 5 minutes balances +// freshness against management load on busy mesh networks. Override it with +// envTunnelCacheTTL when an account needs authorization changes to take effect +// sooner (at the cost of more ValidateTunnelPeer RPCs). const tunnelCacheTTL = 300 * time.Second +// envTunnelCacheTTL overrides tunnelCacheTTL. The value is a Go duration string +// (e.g. "30s", "2m"); an unset, unparseable, or non-positive value keeps the +// default. +const envTunnelCacheTTL = "NB_PROXY_TUNNEL_CACHE_TTL" + // tunnelCachePerAccount caps the number of cached identities per account. // Bounded eviction avoids memory growth in pathological cases (huge peer // roster, brief request bursts) while staying generous for normal use. @@ -60,16 +70,35 @@ type accountBucket struct { order []tunnelCacheKey } -// newTunnelValidationCache constructs a cache with default TTL and bounds. +// newTunnelValidationCache constructs a cache with the configured TTL +// (envTunnelCacheTTL override or default) and default bounds. func newTunnelValidationCache() *tunnelValidationCache { return &tunnelValidationCache{ entries: make(map[types.AccountID]*accountBucket), - ttl: tunnelCacheTTL, + ttl: tunnelCacheTTLFromEnv(), maxSize: tunnelCachePerAccount, now: time.Now, } } +// tunnelCacheTTLFromEnv returns the tunnel-cache TTL, honoring the +// envTunnelCacheTTL override. The override must be a positive Go duration +// string (e.g. "30s", "2m"); anything unset, unparseable, or non-positive +// falls back to tunnelCacheTTL. +func tunnelCacheTTLFromEnv() time.Duration { + raw := strings.TrimSpace(os.Getenv(envTunnelCacheTTL)) + if raw == "" { + return tunnelCacheTTL + } + d, err := time.ParseDuration(raw) + if err != nil || d <= 0 { + log.Warnf("ignoring invalid %s=%q (want a positive Go duration like 30s or 2m); using default %s", + envTunnelCacheTTL, raw, tunnelCacheTTL) + return tunnelCacheTTL + } + return d +} + // get returns a cached response for the key, or nil when missing or // expired. Expired entries are evicted lazily on read. func (c *tunnelValidationCache) get(key tunnelCacheKey) *proto.ValidateTunnelPeerResponse { diff --git a/proxy/internal/auth/tunnel_cache_test.go b/proxy/internal/auth/tunnel_cache_test.go index 1a63dc107..d663b91e7 100644 --- a/proxy/internal/auth/tunnel_cache_test.go +++ b/proxy/internal/auth/tunnel_cache_test.go @@ -169,3 +169,32 @@ func TestTunnelCache_BoundedSizeEvictsOldest(t *testing.T) { assert.NotNil(t, cache.get(keys[1]), "second-newest must remain cached") assert.NotNil(t, cache.get(keys[2]), "newest must remain cached") } + +func TestTunnelCacheTTLFromEnv(t *testing.T) { + t.Run("unset uses default", func(t *testing.T) { + t.Setenv(envTunnelCacheTTL, "") + assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv()) + }) + t.Run("valid duration overrides", func(t *testing.T) { + t.Setenv(envTunnelCacheTTL, "45s") + assert.Equal(t, 45*time.Second, tunnelCacheTTLFromEnv()) + }) + t.Run("whitespace trimmed", func(t *testing.T) { + t.Setenv(envTunnelCacheTTL, " 2m ") + assert.Equal(t, 2*time.Minute, tunnelCacheTTLFromEnv()) + }) + t.Run("unparseable uses default", func(t *testing.T) { + t.Setenv(envTunnelCacheTTL, "nonsense") + assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv()) + }) + t.Run("non-positive uses default", func(t *testing.T) { + t.Setenv(envTunnelCacheTTL, "0s") + assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv()) + t.Setenv(envTunnelCacheTTL, "-30s") + assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv()) + }) + t.Run("constructor honors override", func(t *testing.T) { + t.Setenv(envTunnelCacheTTL, "90s") + assert.Equal(t, 90*time.Second, newTunnelValidationCache().ttl) + }) +} diff --git a/proxy/internal/middleware/builtin/llm_guardrail/factory.go b/proxy/internal/middleware/builtin/llm_guardrail/factory.go index 6dd2a8e8d..9cec0dc68 100644 --- a/proxy/internal/middleware/builtin/llm_guardrail/factory.go +++ b/proxy/internal/middleware/builtin/llm_guardrail/factory.go @@ -10,11 +10,15 @@ import ( ) // Config is the JSON-decoded shape accepted by the factory. The -// runtime path consumes the normalised allowlist; raw config is not +// runtime path consumes the normalised allowlists; raw config is not // retained beyond construction. type Config struct { - ModelAllowlist []string `json:"model_allowlist"` - PromptCapture PromptCapture `json:"prompt_capture"` + // ProviderAllowlists maps a resolved provider id (KeyLLMResolvedProviderID) to + // its model allowlist. A provider present is restricted to those models; one + // absent is unrestricted. Kept per-provider so one provider's list can't leak + // onto another. + ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"` + PromptCapture PromptCapture `json:"prompt_capture"` } // PromptCapture toggles the optional prompt capture + redaction step @@ -54,21 +58,28 @@ func isEmptyJSON(raw []byte) bool { return false } -// normaliseConfig lowercases and trims allowlist entries so the runtime -// match is case-insensitive. Empty entries are dropped. +// normaliseConfig lowercases and trims allowlist entries for case-insensitive +// matching; empty entries drop. A provider whose entries all drop keeps an empty +// (non-nil) list — "deny every model" — distinct from an absent provider +// (unrestricted). func normaliseConfig(cfg Config) Config { - if len(cfg.ModelAllowlist) == 0 { + if len(cfg.ProviderAllowlists) == 0 { + cfg.ProviderAllowlists = nil return cfg } - cleaned := make([]string, 0, len(cfg.ModelAllowlist)) - for _, entry := range cfg.ModelAllowlist { - n := normaliseModel(entry) - if n == "" { - continue + cleaned := make(map[string][]string, len(cfg.ProviderAllowlists)) + for provider, models := range cfg.ProviderAllowlists { + list := make([]string, 0, len(models)) + for _, entry := range models { + n := normaliseModel(entry) + if n == "" { + continue + } + list = append(list, n) } - cleaned = append(cleaned, n) + cleaned[provider] = list } - cfg.ModelAllowlist = cleaned + cfg.ProviderAllowlists = cleaned return cfg } diff --git a/proxy/internal/middleware/builtin/llm_guardrail/middleware.go b/proxy/internal/middleware/builtin/llm_guardrail/middleware.go index eded877ac..1863aff20 100644 --- a/proxy/internal/middleware/builtin/llm_guardrail/middleware.go +++ b/proxy/internal/middleware/builtin/llm_guardrail/middleware.go @@ -83,8 +83,9 @@ func (m *Middleware) MutationsSupported() bool { return false } // prompt capture only affects the metadata emitted alongside an allow. func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) { model, modelPresent := lookupMetadata(in.Metadata, middleware.KeyLLMModel) + providerID, _ := lookupMetadata(in.Metadata, middleware.KeyLLMResolvedProviderID) - if denial := m.evaluateAllowlist(model, modelPresent); denial != nil { + if denial := m.evaluateAllowlist(providerID, model, modelPresent); denial != nil { return denial, nil } @@ -110,20 +111,32 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar // is a no-op. func (m *Middleware) Close() error { return nil } -// evaluateAllowlist returns a deny Output when the configured allowlist -// rejects the model. A nil return means the request should proceed. -func (m *Middleware) evaluateAllowlist(model string, modelPresent bool) *middleware.Output { - if len(m.cfg.ModelAllowlist) == 0 { +// evaluateAllowlist denies when the resolved provider's allowlist rejects the +// model; nil means proceed. Scoped to the provider llm_router resolved, so an +// unrestricted provider (absent from config) is never caught by another's list. +func (m *Middleware) evaluateAllowlist(providerID, model string, modelPresent bool) *middleware.Output { + if len(m.cfg.ProviderAllowlists) == 0 { return nil } - // Fail closed: with an allowlist configured, a request whose model the - // upstream parser could not extract (absent or empty) must be denied rather - // than allowed. This is what enforces the allowlist for URL/path-routed - // providers (Bedrock, Vertex, ...) whose model lives outside the JSON body. + // Restrictions exist but the resolved provider is unknown, so we can't tell + // if this request targets a restricted provider — fail closed. llm_router + // normally stamps the provider first, so this is a defensive guard. + if providerID == "" { + return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown) + } + allowlist, restricted := m.cfg.ProviderAllowlists[providerID] + if !restricted { + // This provider has no allowlist (some authorising policy left it + // unrestricted); management owns any per-policy/group decision. + return nil + } + // Fail closed: with an allowlist in effect for this provider, a request whose + // model the parser couldn't extract (absent/empty) is denied. This enforces + // the allowlist for path-routed providers (Bedrock, Vertex) with no body model. if !modelPresent || normaliseModel(model) == "" { return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown) } - if m.modelInAllowlist(model) { + if modelInAllowlist(allowlist, model) { return nil } return denyModel(model, denyCodeModel, denyMessageModel, denyReasonModel) @@ -151,14 +164,15 @@ func denyModel(model, code, message, reason string) *middleware.Output { } } -// modelInAllowlist reports whether the model matches any allowlist -// entry under the case-insensitive, trim-tolerant comparison rule. -func (m *Middleware) modelInAllowlist(model string) bool { +// modelInAllowlist reports whether the model matches any entry in the supplied +// (already-normalised) allowlist under the case-insensitive, trim-tolerant +// comparison rule. +func modelInAllowlist(allowlist []string, model string) bool { normalised := normaliseModel(model) if normalised == "" { return false } - for _, allowed := range m.cfg.ModelAllowlist { + for _, allowed := range allowlist { if allowed == normalised { return true } diff --git a/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go b/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go index cd7e256dd..5f35fefd3 100644 --- a/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go +++ b/proxy/internal/middleware/builtin/llm_guardrail/middleware_test.go @@ -26,6 +26,25 @@ func newInput(meta ...middleware.KV) *middleware.Input { return &middleware.Input{Slot: middleware.SlotOnRequest, Metadata: meta} } +const ( + testProvider = "prov-1" + otherProvider = "prov-2" +) + +// providerCfg builds a Config restricting testProvider to the given models. +func providerCfg(models ...string) Config { + return Config{ProviderAllowlists: map[string][]string{testProvider: models}} +} + +// newInputProvider builds an input that carries a resolved provider id (as +// llm_router would stamp) plus any extra metadata. +func newInputProvider(provider string, meta ...middleware.KV) *middleware.Input { + all := make([]middleware.KV, 0, len(meta)+1) + all = append(all, middleware.KV{Key: middleware.KeyLLMResolvedProviderID, Value: provider}) + all = append(all, meta...) + return &middleware.Input{Slot: middleware.SlotOnRequest, Metadata: all} +} + func TestMiddlewareIdentity(t *testing.T) { mw := New(Config{}) assert.Equal(t, ID, mw.ID(), "middleware ID must be llm_guardrail") @@ -47,12 +66,12 @@ func TestMiddlewareIdentity(t *testing.T) { func TestAllowlistEmptyAllowsAnyModel(t *testing.T) { mw := New(Config{}) - out, err := mw.Invoke(context.Background(), newInput( + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, )) require.NoError(t, err) require.NotNil(t, out) - assert.Equal(t, middleware.DecisionAllow, out.Decision, "empty allowlist must allow any model") + assert.Equal(t, middleware.DecisionAllow, out.Decision, "no provider allowlists must allow any model") v, ok := metaValue(t, out.Metadata, middleware.KeyLLMPolicyDecision) require.True(t, ok, "decision metadata must be emitted") assert.Equal(t, "allow", v, "decision must be allow") @@ -62,8 +81,8 @@ func TestAllowlistEmptyAllowsAnyModel(t *testing.T) { } func TestAllowlistMatchAllows(t *testing.T) { - mw := New(Config{ModelAllowlist: []string{"gpt-4o", "claude-opus-4"}}) - out, err := mw.Invoke(context.Background(), newInput( + mw := New(providerCfg("gpt-4o", "claude-opus-4")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, )) require.NoError(t, err) @@ -71,8 +90,8 @@ func TestAllowlistMatchAllows(t *testing.T) { } func TestAllowlistMissDenies(t *testing.T) { - mw := New(Config{ModelAllowlist: []string{"gpt-4o"}}) - out, err := mw.Invoke(context.Background(), newInput( + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"}, )) require.NoError(t, err) @@ -91,10 +110,10 @@ func TestAllowlistMissDenies(t *testing.T) { } func TestAllowlistCaseInsensitive(t *testing.T) { - mw := New(Config{ModelAllowlist: []string{" GPT-4o ", "Claude-OPUS-4"}}) + mw := New(providerCfg(" GPT-4o ", "Claude-OPUS-4")) cases := []string{"gpt-4o", "GPT-4O", " claude-opus-4 "} for _, model := range cases { - out, err := mw.Invoke(context.Background(), newInput( + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: model}, )) require.NoError(t, err) @@ -103,14 +122,15 @@ func TestAllowlistCaseInsensitive(t *testing.T) { } func TestAllowlistMissingModelKeyDenies(t *testing.T) { - // Fail closed: with an allowlist configured, a request whose model the - // parser could not extract (URL/path-routed providers such as Bedrock or - // Vertex whose shape wasn't recognised) must be denied, not allowed. - mw := New(Config{ModelAllowlist: []string{"gpt-4o"}}) - out, err := mw.Invoke(context.Background(), newInput()) + // Fail closed: with an allowlist in effect for the resolved provider, a + // request whose model the parser could not extract (URL/path-routed + // providers such as Bedrock or Vertex whose shape wasn't recognised) must be + // denied, not allowed. + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider)) require.NoError(t, err) require.NotNil(t, out) - assert.Equal(t, middleware.DecisionDeny, out.Decision, "absent model must be denied when an allowlist is set") + assert.Equal(t, middleware.DecisionDeny, out.Decision, "absent model must be denied when the provider is restricted") assert.Equal(t, 403, out.DenyStatus, "deny status must be 403") require.NotNil(t, out.DenyReason, "deny reason must be populated") assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown") @@ -122,26 +142,101 @@ func TestAllowlistMissingModelKeyDenies(t *testing.T) { func TestAllowlistEmptyModelValueDenies(t *testing.T) { // A present-but-empty model is as undeterminable as an absent one. - mw := New(Config{ModelAllowlist: []string{"gpt-4o"}}) - out, err := mw.Invoke(context.Background(), newInput( + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: " "}, )) require.NoError(t, err) require.NotNil(t, out) - assert.Equal(t, middleware.DecisionDeny, out.Decision, "empty model must be denied when an allowlist is set") + assert.Equal(t, middleware.DecisionDeny, out.Decision, "empty model must be denied when the provider is restricted") require.NotNil(t, out.DenyReason, "deny reason must be populated") assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown") } func TestAllowlistEmptyListAllowsMissingModel(t *testing.T) { - // Without an allowlist there is nothing to enforce, so a missing model is - // still allowed — the fail-closed rule only applies when a list is set. + // Without any provider allowlists there is nothing to enforce, so a missing + // model is still allowed — the fail-closed rule only applies when a + // restriction is in effect. mw := New(Config{}) out, err := mw.Invoke(context.Background(), newInput()) require.NoError(t, err) assert.Equal(t, middleware.DecisionAllow, out.Decision, "no allowlist must allow even without a model") } +func TestUnrestrictedProviderAllowsAnyModel(t *testing.T) { + // The request resolved to otherProvider, which has no allowlist, so its + // traffic must not be caught by testProvider's restriction — the + // cross-provider-leak / false-deny guard. + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInputProvider(otherProvider, + middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"}, + )) + require.NoError(t, err) + assert.Equal(t, middleware.DecisionAllow, out.Decision, "an unrestricted provider must not inherit another provider's allowlist") +} + +func TestPerProviderAllowlistsAreIsolated(t *testing.T) { + // gpt-4o is allowed only on testProvider; claude-opus-4 only on + // otherProvider. A model allowlisted for one provider must not be usable on + // the other — the fail-closed layer never unions allowlists across providers. + mw := New(Config{ProviderAllowlists: map[string][]string{ + testProvider: {"gpt-4o"}, + otherProvider: {"claude-opus-4"}, + }}) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, + middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-opus-4"}, + )) + require.NoError(t, err) + assert.Equal(t, middleware.DecisionDeny, out.Decision, "claude-opus-4 is allowed only on otherProvider, not testProvider") + require.NotNil(t, out.DenyReason) + assert.Equal(t, "llm_policy.model_blocked", out.DenyReason.Code, "cross-provider model must be blocked, not model_unknown") +} + +func TestRestrictionsButNoResolvedProviderFailsClosed(t *testing.T) { + // Restrictions exist for the account but the resolved provider id is absent, + // so the request cannot be scoped to a provider. Fail closed rather than + // wave it through. + mw := New(providerCfg("gpt-4o")) + out, err := mw.Invoke(context.Background(), newInput( + middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + )) + require.NoError(t, err) + require.NotNil(t, out) + assert.Equal(t, middleware.DecisionDeny, out.Decision, "missing resolved provider must fail closed when restrictions exist") + require.NotNil(t, out.DenyReason) + assert.Equal(t, "llm_policy.model_unknown", out.DenyReason.Code, "deny code must be model_unknown") +} + +func TestEnabledButEmptyAllowlistDeniesEveryModel(t *testing.T) { + // An allowlist-enabled provider with zero models is distinct from an + // unrestricted (absent) provider: it must deny every model. + mw := New(providerCfg()) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, + middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + )) + require.NoError(t, err) + require.NotNil(t, out) + assert.Equal(t, middleware.DecisionDeny, out.Decision, "an enabled-but-empty allowlist must deny every model") + require.NotNil(t, out.DenyReason) + assert.Equal(t, "llm_policy.model_blocked", out.DenyReason.Code, "deny code must be model_blocked, not model_unknown") +} + +func TestFactoryAllEmptyEntriesDenyEveryModel(t *testing.T) { + // All the provider's entries are blank; they collapse to a non-nil empty + // list (deny everything for that provider), not "no restriction". + raw := []byte(`{"provider_allowlists":{"prov-1":[""," "]}}`) + mw, err := Factory{}.New(raw) + require.NoError(t, err) + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, + middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + )) + require.NoError(t, err) + require.NotNil(t, out) + assert.Equal(t, middleware.DecisionDeny, out.Decision, "all-blank allowlist entries must still restrict the provider") + require.NotNil(t, out.DenyReason) + assert.Equal(t, "llm_policy.model_blocked", out.DenyReason.Code, "deny code must be model_blocked") +} + func TestPromptCaptureDisabledEmitsNoPrompt(t *testing.T) { mw := New(Config{}) out, err := mw.Invoke(context.Background(), newInput( @@ -217,8 +312,8 @@ func TestFactoryAcceptsZeroConfigs(t *testing.T) { func TestFactoryDecodesValidConfig(t *testing.T) { cfg := Config{ - ModelAllowlist: []string{"gpt-4o"}, - PromptCapture: PromptCapture{Enabled: true, RedactPii: true}, + ProviderAllowlists: map[string][]string{testProvider: {"gpt-4o"}}, + PromptCapture: PromptCapture{Enabled: true, RedactPii: true}, } raw, err := json.Marshal(cfg) require.NoError(t, err, "marshalling test config must succeed") @@ -234,15 +329,15 @@ func TestFactoryRejectsMalformedJSON(t *testing.T) { } func TestFactoryNormalisesAllowlist(t *testing.T) { - raw := []byte(`{"model_allowlist":[" GPT-4o ","",""," Claude-3 "]}`) + raw := []byte(`{"provider_allowlists":{"prov-1":[" GPT-4o ","",""," Claude-3 "]}}`) mw, err := Factory{}.New(raw) require.NoError(t, err) - out, err := mw.Invoke(context.Background(), newInput( + out, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "gpt-4o"}, )) require.NoError(t, err) assert.Equal(t, middleware.DecisionAllow, out.Decision, "factory must lowercase + trim allowlist entries") - out2, err := mw.Invoke(context.Background(), newInput( + out2, err := mw.Invoke(context.Background(), newInputProvider(testProvider, middleware.KV{Key: middleware.KeyLLMModel, Value: "claude-3"}, )) require.NoError(t, err) diff --git a/proxy/internal/middleware/builtin/llm_limit_check/middleware.go b/proxy/internal/middleware/builtin/llm_limit_check/middleware.go index bebe4dca4..42ac56b9b 100644 --- a/proxy/internal/middleware/builtin/llm_limit_check/middleware.go +++ b/proxy/internal/middleware/builtin/llm_limit_check/middleware.go @@ -175,7 +175,7 @@ func denyFromManagement(resp *proto.CheckLLMPolicyLimitsResponse) *middleware.Ou DenyStatus: 403, DenyReason: &middleware.DenyReason{ Code: code, - Message: "LLM policy limit exceeded", + Message: denyMessageForCode(code), }, Metadata: []middleware.KV{ {Key: middleware.KeyLLMPolicyDecision, Value: "deny"}, @@ -184,6 +184,21 @@ func denyFromManagement(resp *proto.CheckLLMPolicyLimitsResponse) *middleware.Ou } } +// denyMessageForCode maps a management deny code to a public message. +// Model-allowlist rejections get a model-specific message matching the +// local guardrail; everything else keeps the generic quota wording. The +// message stays generic so it never leaks internal quota detail. +func denyMessageForCode(code string) string { + switch code { + case "llm_policy.model_blocked": + return "model is not in the policy allowlist" + case "llm_policy.model_unknown": + return "request model could not be determined for the policy allowlist" + default: + return "LLM policy limit exceeded" + } +} + // lookupKV returns the value associated with key, or the empty // string when absent. func lookupKV(kvs []middleware.KV, key string) string { diff --git a/proxy/internal/middleware/builtin/llm_limit_check/middleware_test.go b/proxy/internal/middleware/builtin/llm_limit_check/middleware_test.go index 2c26c2abe..87aa8e9e9 100644 --- a/proxy/internal/middleware/builtin/llm_limit_check/middleware_test.go +++ b/proxy/internal/middleware/builtin/llm_limit_check/middleware_test.go @@ -115,6 +115,46 @@ func TestInvoke_DenyConvertsToProxyDeny(t *testing.T) { assert.NotContains(t, out.DenyReason.Message, "1000", "internal cap numbers must not reach the caller") } +// TestInvoke_ModelDenyMessages proves a model-allowlist rejection gets a +// model-specific public message rather than the generic quota wording, so a +// blocked or undetermined model reads consistently with the local guardrail. +func TestInvoke_ModelDenyMessages(t *testing.T) { + cases := []struct { + name string + code string + message string + }{ + {"blocked", "llm_policy.model_blocked", "model is not in the policy allowlist"}, + {"unknown", "llm_policy.model_unknown", "request model could not be determined for the policy allowlist"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + mgmt := &fakeMgmt{ + checkResp: &proto.CheckLLMPolicyLimitsResponse{ + Decision: "deny", + DenyCode: tc.code, + }, + } + m := New(mgmt, nil) + + out := runInvoke(t, m, &middleware.Input{ + AccountID: "acc-1", + UserGroups: []string{"grp-engineers"}, + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMResolvedProviderID, Value: "prov-1"}, + {Key: middleware.KeyLLMModel, Value: "some-model"}, + }, + }) + + assert.Equal(t, middleware.DecisionDeny, out.Decision) + require.NotNil(t, out.DenyReason, "deny envelope must carry a reason payload") + assert.Equal(t, tc.code, out.DenyReason.Code, "canonical deny code surfaces to the caller") + assert.Equal(t, tc.message, out.DenyReason.Message, + "model denials must use a model-specific message, matching the local guardrail") + }) + } +} + // TestInvoke_NoMgmtClientPassesThrough proves the partial-wiring // safety: a middleware constructed without a management client // allows every request without attribution. This makes a half-set-up diff --git a/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go b/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go index 0074411cc..1eae26d73 100644 --- a/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go +++ b/proxy/internal/middleware/builtin/llm_request_parser/guardrail_allowlist_test.go @@ -25,10 +25,17 @@ func runParserGuardrail(t *testing.T, url string, body []byte, allowlist []strin }) require.NoError(t, err, "parser must not error") - guard := llm_guardrail.New(llm_guardrail.Config{ModelAllowlist: allowlist}) + const providerID = "prov-under-test" + guard := llm_guardrail.New(llm_guardrail.Config{ + ProviderAllowlists: map[string][]string{providerID: allowlist}, + }) + // The real chain has llm_router stamp the resolved provider id before the + // guardrail runs; the parser doesn't, so add it here so the guardrail can + // scope the allowlist to this provider. + meta := append([]middleware.KV{{Key: middleware.KeyLLMResolvedProviderID, Value: providerID}}, parsed.Metadata...) out, err := guard.Invoke(context.Background(), &middleware.Input{ Slot: middleware.SlotOnRequest, - Metadata: parsed.Metadata, + Metadata: meta, }) require.NoError(t, err, "guardrail must not error") require.NotNil(t, out, "guardrail must return an output") From e3c41281641f5840408f73d852fb7c7914cbbac6 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 28 Jul 2026 07:09:16 +0200 Subject: [PATCH 06/32] [client] Exit GUI immediately on Windows session end (#6878) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Wails v3 runs the full app teardown synchronously inside WM_ENDSESSION, overrunning the 5s end-session budget and triggering the "app is preventing shutdown" screen with a forced kill. Intercept WM_QUERYENDSESSION/WM_ENDSESSION and exit at once instead, and suppress error dialogs, toasts, and the hide-on-close hooks once shutdown or a tray quit has begun. ## Describe your changes ## Issue ticket number and link ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **Bug Fixes** * Improved shutdown handling on Windows by properly responding to system end-session requests. * Prevented the main window, settings window, and error dialogs from reopening or interfering while the app is closing. * Updated tray “Quit” flow to begin shutdown immediately, ensuring active profile/connection operations complete cleanly. * Suppressed UI notifications during shutdown to avoid stray messages after exit starts. --- client/ui/main.go | 6 +++++ client/ui/services/shutdown.go | 24 +++++++++++++++++++ client/ui/services/windowmanager.go | 6 +++++ client/ui/shutdown_other.go | 7 ++++++ client/ui/shutdown_windows.go | 36 +++++++++++++++++++++++++++++ client/ui/tray.go | 1 + client/ui/tray_notify.go | 3 +++ 7 files changed, 83 insertions(+) create mode 100644 client/ui/services/shutdown.go create mode 100644 client/ui/shutdown_other.go create mode 100644 client/ui/shutdown_windows.go diff --git a/client/ui/main.go b/client/ui/main.go index 4889bad79..9a3e17743 100644 --- a/client/ui/main.go +++ b/client/ui/main.go @@ -281,6 +281,9 @@ func newApplication(onSecondInstance func()) *application.App { Linux: application.LinuxOptions{ ProgramName: "netbird", }, + Windows: application.WindowsOptions{ + WndProcInterceptor: endSessionInterceptor(), + }, SingleInstance: &application.SingleInstanceOptions{ UniqueID: "io.netbird.ui", OnSecondInstanceLaunch: func(_ application.SecondInstanceData) { @@ -367,6 +370,9 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat // Hide instead of quit on close; "really quit" is reached via tray -> Quit. window.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) { + if services.ShuttingDown() { + return + } e.Cancel() window.Hide() }) diff --git a/client/ui/services/shutdown.go b/client/ui/services/shutdown.go new file mode 100644 index 000000000..0da51940c --- /dev/null +++ b/client/ui/services/shutdown.go @@ -0,0 +1,24 @@ +package services + +import "sync/atomic" + +var ( + sessionEnding atomic.Bool + quitting atomic.Bool +) + +func BeginSessionEnd() { + sessionEnding.Store(true) +} + +func AbortSessionEnd() { + sessionEnding.Store(false) +} + +func BeginShutdown() { + quitting.Store(true) +} + +func ShuttingDown() bool { + return sessionEnding.Load() || quitting.Load() +} diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index 1185ec729..bac9790b4 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -154,6 +154,9 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo }) // Hide (not destroy) on close to keep React state; reset to General for a flash-free reopen. s.settings.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) { + if ShuttingDown() { + return + } e.Cancel() s.app.Event.Emit(EventSettingsOpen, "general") s.settings.Hide() @@ -393,6 +396,9 @@ func (s *WindowManager) CloseWelcome() { // OpenError shows the custom error dialog; title/message are pre-localised and ride in the // start URL. A second error replaces the open one via SetURL. Singleton, destroyed on close. func (s *WindowManager) OpenError(title, message string) { + if ShuttingDown() { + return + } s.mu.Lock() defer s.mu.Unlock() startURL := errorDialogURL(title, message) diff --git a/client/ui/shutdown_other.go b/client/ui/shutdown_other.go new file mode 100644 index 000000000..6e617233f --- /dev/null +++ b/client/ui/shutdown_other.go @@ -0,0 +1,7 @@ +//go:build !windows && !android && !ios && !freebsd && !js + +package main + +func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) { + return nil +} diff --git a/client/ui/shutdown_windows.go b/client/ui/shutdown_windows.go new file mode 100644 index 000000000..fbb92a518 --- /dev/null +++ b/client/ui/shutdown_windows.go @@ -0,0 +1,36 @@ +//go:build windows + +package main + +import ( + "os" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/ui/services" +) + +const ( + wmQueryEndSession = 0x0011 + wmEndSession = 0x0016 +) + +func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) { + return func(_ uintptr, msg uint32, wParam, _ uintptr) (uintptr, bool) { + switch msg { + case wmQueryEndSession: + services.BeginSessionEnd() + return 1, true + case wmEndSession: + if wParam == 0 { + services.AbortSessionEnd() + return 0, true + } + log.Info("windows session is ending; exiting immediately") + os.Exit(0) + return 0, true + default: + return 0, false + } + } +} diff --git a/client/ui/tray.go b/client/ui/tray.go index 49e883f89..3050d159a 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -452,6 +452,7 @@ func (t *Tray) buildMenu() *application.Menu { } func (t *Tray) handleQuit() { + services.BeginShutdown() t.profileMu.Lock() if t.switchCancel != nil { t.switchCancel() diff --git a/client/ui/tray_notify.go b/client/ui/tray_notify.go index 6c60e3d4b..d1117b57b 100644 --- a/client/ui/tray_notify.go +++ b/client/ui/tray_notify.go @@ -25,6 +25,9 @@ type sendFn func(notifications.NotificationOptions) error // event-dispatch goroutine that panic is fatal process-wide; recover() turns // it into a logged no-op. func safeSendNotification(send sendFn, what string, opts notifications.NotificationOptions) (err error) { + if services.ShuttingDown() { + return nil + } defer func() { if r := recover(); r != nil { log.Errorf("notify %s: recovered from panic (notification bus unavailable): %v", what, r) From 2f268c814187eda5af3617644f52d07cd3e72a66 Mon Sep 17 00:00:00 2001 From: Riccardo Manfrin <3090891+riccardomanfrin@users.noreply.github.com> Date: Tue, 28 Jul 2026 07:10:58 +0200 Subject: [PATCH 07/32] [client] fix stale routing peer on overlapping-prefix network removal (#6799) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes The AllowedIPs reference counter ([refcounter/types.go#L9](https://github.com/netbirdio/netbird/blob/e1a24376a/client/internal/routemanager/refcounter/types.go#L9)) was keyed only by prefix and stored a single active peer set by the first registrar, never swapped. When two networks advertised the same prefix via different routing peers, removing the one whose peer was installed in WireGuard left the prefix pointing at the removed peer instead of the surviving one — traffic kept flowing to the old peer until a manual `netbird down/up`. Made the AllowedIPs counter peer-aware: it tracks a per-peer reference count per prefix plus the installed peer, and swaps WireGuard to a surviving peer when the active one releases its last reference (removes the prefix when none remain). `Decrement` now takes the peer key so the exact incremented peer is released; the static handler records its selected routing peer like the dynamic and DNS handlers already did. The generic `Counter` (routes, exclusion, ipset) is unchanged. ## Issue ticket number and link No public issue — reported internally (routes not updating without `netbird down/up` when two networks share a subnet). Root cause is the prefix-only key at [refcounter/types.go#L9](https://github.com/netbirdio/netbird/blob/e1a24376a/client/internal/routemanager/refcounter/types.go#L9). ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [x] Created tests that fail without the change (if possible) > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) Internal client-side routing fix. No public API, CLI, or config change — only the WireGuard AllowedIPs hand-off when overlapping-prefix networks are removed. ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: N/A --- View with Codesmith Autofix with Codesmith Need help on this PR? Tag /codesmith with what you need. Autofix is disabled. ## Summary by CodeRabbit * **Bug Fixes** * Improved routing behavior when multiple peers share the same Allowed IP by making Allowed IP reference tracking peer-aware. * Allowed IPs now correctly decrement using the active peer key and transfer to another surviving active peer when the current peer is removed. * Prevented stale routing and incorrect reference cleanup during route and DNS-driven teardown. * **Tests** * Added/extended coverage for peer handoffs, repeated references, non-active peer removal, flushing behavior, and self-healing after swap add/remove failures. --- .../routemanager/dnsinterceptor/handler.go | 6 +- client/internal/routemanager/dynamic/route.go | 4 +- client/internal/routemanager/manager.go | 2 +- .../routemanager/refcounter/allowedips.go | 185 ++++++++++++++ .../refcounter/allowedips_test.go | 241 ++++++++++++++++++ .../internal/routemanager/refcounter/types.go | 6 +- client/internal/routemanager/static/route.go | 14 +- 7 files changed, 447 insertions(+), 11 deletions(-) create mode 100644 client/internal/routemanager/refcounter/allowedips.go create mode 100644 client/internal/routemanager/refcounter/allowedips_test.go diff --git a/client/internal/routemanager/dnsinterceptor/handler.go b/client/internal/routemanager/dnsinterceptor/handler.go index b784cc274..f92300bfd 100644 --- a/client/internal/routemanager/dnsinterceptor/handler.go +++ b/client/internal/routemanager/dnsinterceptor/handler.go @@ -95,7 +95,7 @@ func (d *DnsInterceptor) RemoveRoute() error { // AllowedIPs should use real IPs if d.currentPeerKey != "" { - if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err)) } } @@ -172,7 +172,7 @@ func (d *DnsInterceptor) removeAllowedIP(realPrefix netip.Prefix) error { } // AllowedIPs use real IPs - if _, err := d.allowedIPsRefcounter.Decrement(realPrefix); err != nil { + if _, err := d.allowedIPsRefcounter.Decrement(realPrefix, d.currentPeerKey); err != nil { return fmt.Errorf("remove allowed IP %s: %v", realPrefix, err) } @@ -205,7 +205,7 @@ func (d *DnsInterceptor) RemoveAllowedIPs() error { for _, prefixes := range d.interceptedDomains { for _, prefix := range prefixes { // AllowedIPs use real IPs - if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err)) } } diff --git a/client/internal/routemanager/dynamic/route.go b/client/internal/routemanager/dynamic/route.go index 3fe8a4bb3..bb3b1c59c 100644 --- a/client/internal/routemanager/dynamic/route.go +++ b/client/internal/routemanager/dynamic/route.go @@ -135,7 +135,7 @@ func (r *Route) RemoveAllowedIPs() error { var merr *multierror.Error for _, domainPrefixes := range r.dynamicDomains { for _, prefix := range domainPrefixes { - if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err)) } } @@ -320,7 +320,7 @@ func (r *Route) removeRoutes(prefixes []netip.Prefix) ([]netip.Prefix, error) { merr = multierror.Append(merr, fmt.Errorf("remove dynamic route for IP %s: %w", prefix, err)) } if r.currentPeerKey != "" { - if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil { + if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil { merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err)) } } diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go index ef69b81a4..7a818b539 100644 --- a/client/internal/routemanager/manager.go +++ b/client/internal/routemanager/manager.go @@ -216,7 +216,7 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) { ) } - m.allowedIPsRefCounter = refcounter.New( + m.allowedIPsRefCounter = refcounter.NewAllowedIPs( func(prefix netip.Prefix, peerKey string) (string, error) { // save peerKey to use it in the remove function return peerKey, m.wgInterface.AddAllowedIP(peerKey, prefix) diff --git a/client/internal/routemanager/refcounter/allowedips.go b/client/internal/routemanager/refcounter/allowedips.go new file mode 100644 index 000000000..6f23162f6 --- /dev/null +++ b/client/internal/routemanager/refcounter/allowedips.go @@ -0,0 +1,185 @@ +package refcounter + +import ( + "errors" + "fmt" + "net/netip" + "sort" + "sync" + + "github.com/hashicorp/go-multierror" + + nberrors "github.com/netbirdio/netbird/client/errors" +) + +// allowedIPsEntry holds the per-peer reference counts for a single prefix and which peer is +// currently installed in WireGuard. WireGuard allows a prefix on exactly one peer, so at most +// one peer is active at a time even when several peers reference the prefix. +type allowedIPsEntry struct { + // peers maps a peerKey to the number of references holding the prefix for that peer. + peers map[string]int + // active is the peerKey currently installed in WireGuard for this prefix ("" if none). + active string + // total is the sum of all per-peer reference counts (kept in sync with peers). + total int +} + +// AllowedIPsRefCounter is a peer-aware reference counter for WireGuard AllowedIPs. +// +// The generic Counter keys only by prefix and remembers a single Out value set by the first +// caller, which it never changes. That is wrong for AllowedIPs: two independent watchers (or +// multiple resolved domains) can reference the same prefix through different peers, and when the +// peer currently installed in WireGuard releases its last reference the prefix must be handed over +// to a surviving peer instead of being left pointing at the released one. +// +// It calls add/remove (which program WireGuard) only on the transitions that matter: +// - add on the first reference for a prefix, or when swapping the active peer; +// - remove on the last reference for a prefix, or on the old peer during a swap. +type AllowedIPsRefCounter struct { + mu sync.Mutex + entries map[netip.Prefix]*allowedIPsEntry + add AddFunc[netip.Prefix, string, string] + remove RemoveFunc[netip.Prefix, string] +} + +// NewAllowedIPs creates a new peer-aware AllowedIPs reference counter. +// add programs a prefix on a peer in WireGuard and returns the peerKey to store as the active peer. +// remove unprograms the prefix from the given peer. +func NewAllowedIPs(add AddFunc[netip.Prefix, string, string], remove RemoveFunc[netip.Prefix, string]) *AllowedIPsRefCounter { + return &AllowedIPsRefCounter{ + entries: map[netip.Prefix]*allowedIPsEntry{}, + add: add, + remove: remove, + } +} + +// Increment adds a reference to prefix for peerKey. WireGuard is programmed only for the first +// reference to a prefix; while a different peer is already installed the prefix is left with it +// (first peer wins, HA at the WireGuard layer is not possible) and only the reference count is kept. +func (rm *AllowedIPsRefCounter) Increment(prefix netip.Prefix, peerKey string) (Ref[string], error) { + rm.mu.Lock() + defer rm.mu.Unlock() + + e, ok := rm.entries[prefix] + if !ok { + e = &allowedIPsEntry{peers: map[string]int{}} + rm.entries[prefix] = e + } + + logCallerF("Increasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]", + prefix, peerKey, e.peers[peerKey], e.peers[peerKey]+1, e.total, e.total+1, e.active) + + // Program WireGuard only when nothing is installed yet for this prefix. + if e.active == "" { + out, err := rm.add(prefix, peerKey) + if errors.Is(err, ErrIgnore) { + if e.total == 0 { + delete(rm.entries, prefix) + } + return Ref[string]{Count: e.total, Out: e.active}, nil + } + if err != nil { + if e.total == 0 { + delete(rm.entries, prefix) + } + return Ref[string]{}, fmt.Errorf("failed to add allowed IP %v for peer %s: %w", prefix, peerKey, err) + } + e.active = out + } + + e.peers[peerKey]++ + e.total++ + + return Ref[string]{Count: e.total, Out: e.active}, nil +} + +// Decrement removes a reference to prefix for peerKey. When the peer currently installed in +// WireGuard releases its last reference, the prefix is swapped to a surviving peer if one exists, +// otherwise it is removed from WireGuard. +func (rm *AllowedIPsRefCounter) Decrement(prefix netip.Prefix, peerKey string) (Ref[string], error) { + rm.mu.Lock() + defer rm.mu.Unlock() + + e, ok := rm.entries[prefix] + if !ok { + logCallerF("No allowed IP reference found for prefix %v", prefix) + return Ref[string]{}, nil + } + + if e.peers[peerKey] > 0 { + logCallerF("Decreasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]", + prefix, peerKey, e.peers[peerKey], e.peers[peerKey]-1, e.total, e.total-1, e.active) + e.peers[peerKey]-- + e.total-- + if e.peers[peerKey] == 0 { + delete(e.peers, peerKey) + } + } else { + logCallerF("No allowed IP reference found for prefix %v peer %s", prefix, peerKey) + } + + // If the peer currently installed in WireGuard still holds references, nothing to reprogram. + // Keying the check on the active peer (not the one just released) makes this self-healing: + // a prior swap whose remove/add failed leaves e.active pointing at a peer with no references, + // and this retries the hand-off on the next Decrement instead of getting stuck. + if e.active != "" && e.peers[e.active] > 0 { + return Ref[string]{Count: e.total, Out: e.active}, nil + } + + // Detach the stale/gone active peer from WireGuard before reprogramming. + if e.active != "" { + if err := rm.remove(prefix, e.active); err != nil { + return Ref[string]{Count: e.total, Out: e.active}, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err) + } + e.active = "" + } + + // Hand the prefix over to a surviving peer, or drop the entry when none remain. + if survivor, ok := pickSurvivor(e.peers); ok { + out, err := rm.add(prefix, survivor) + if err != nil { + return Ref[string]{Count: e.total, Out: ""}, fmt.Errorf("swap allowed IP %v to peer %s: %w", prefix, survivor, err) + } + e.active = out + return Ref[string]{Count: e.total, Out: e.active}, nil + } + + delete(rm.entries, prefix) + return Ref[string]{Count: 0, Out: ""}, nil +} + +// Flush removes all prefixes from WireGuard and clears the counter. +func (rm *AllowedIPsRefCounter) Flush() error { + rm.mu.Lock() + defer rm.mu.Unlock() + + var merr *multierror.Error + for prefix, e := range rm.entries { + if e.active == "" { + continue + } + logCallerF("Flushing allowed IP for prefix %v peer %s", prefix, e.active) + if err := rm.remove(prefix, e.active); err != nil { + merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err)) + } + } + + clear(rm.entries) + + return nberrors.FormatErrorOrNil(merr) +} + +// pickSurvivor deterministically selects a peer still referencing the prefix. WireGuard cannot do +// multipath for a single prefix, so any surviving peer is a valid winner; the choice is made stable +// (lowest peerKey) for predictable behavior and testability. +func pickSurvivor(peers map[string]int) (string, bool) { + if len(peers) == 0 { + return "", false + } + keys := make([]string, 0, len(peers)) + for k := range peers { + keys = append(keys, k) + } + sort.Strings(keys) + return keys[0], true +} diff --git a/client/internal/routemanager/refcounter/allowedips_test.go b/client/internal/routemanager/refcounter/allowedips_test.go new file mode 100644 index 000000000..835142083 --- /dev/null +++ b/client/internal/routemanager/refcounter/allowedips_test.go @@ -0,0 +1,241 @@ +package refcounter + +import ( + "errors" + "net/netip" + "testing" +) + +// fakeWG models WireGuard's cryptokey routing: a prefix can be installed on exactly one peer. +// failAdd/failRemove make the next add/remove fail once, to exercise the self-healing error paths. +type fakeWG struct { + installed map[netip.Prefix]string + adds int + removes int + failAdd bool + failRemove bool +} + +func newFakeWG() *fakeWG { + return &fakeWG{installed: map[netip.Prefix]string{}} +} + +func (f *fakeWG) counter() *AllowedIPsRefCounter { + return NewAllowedIPs( + func(prefix netip.Prefix, peerKey string) (string, error) { + if f.failAdd { + f.failAdd = false + return "", errors.New("add failed") + } + f.adds++ + f.installed[prefix] = peerKey + return peerKey, nil + }, + func(prefix netip.Prefix, peerKey string) error { + if f.failRemove { + f.failRemove = false + return errors.New("remove failed") + } + f.removes++ + // only clear if this peer is the one installed, mirroring wg semantics + if f.installed[prefix] == peerKey { + delete(f.installed, prefix) + } + return nil + }, + ) +} + +func mustPrefix(t *testing.T, s string) netip.Prefix { + t.Helper() + p, err := netip.ParsePrefix(s) + if err != nil { + t.Fatalf("parse prefix %q: %v", s, err) + } + return p +} + +func mustIncrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] { + t.Helper() + ref, err := c.Increment(p, peer) + if err != nil { + t.Fatalf("Increment(%v, %s): %v", p, peer, err) + } + return ref +} + +func mustDecrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] { + t.Helper() + ref, err := c.Decrement(p, peer) + if err != nil { + t.Fatalf("Decrement(%v, %s): %v", p, peer, err) + } + return ref +} + +// TestAllowedIPs_SwapOnActivePeerRemoval reproduces the reported bug: two networks with the same +// prefix routed by different peers. Removing the network whose peer is installed must hand the +// prefix over to the surviving peer instead of leaving it on the removed one. +func TestAllowedIPs_SwapOnActivePeerRemoval(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + // First peer wins while both are present. + if got := f.installed[p]; got != "peerA" { + t.Fatalf("expected peerA installed, got %q", got) + } + + // Remove the active peer's network -> must swap to peerB. + mustDecrement(t, c, p, "peerA") + if got := f.installed[p]; got != "peerB" { + t.Fatalf("BUG: prefix stuck on removed peer, want peerB got %q", got) + } + + // Remove the last one -> prefix gone. + mustDecrement(t, c, p, "peerB") + if _, ok := f.installed[p]; ok { + t.Fatalf("expected prefix removed, still installed on %q", f.installed[p]) + } +} + +// TestAllowedIPs_RemoveNonActivePeer removing a non-installed peer must not touch WireGuard. +func TestAllowedIPs_RemoveNonActivePeer(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + removesBefore := f.removes + + mustDecrement(t, c, p, "peerB") + if f.installed[p] != "peerA" { + t.Fatalf("active peer must stay peerA, got %q", f.installed[p]) + } + if f.removes != removesBefore { + t.Fatalf("removing a non-active peer must not call wg remove") + } +} + +// TestAllowedIPs_SamePeerMultipleRefs two references via the same peer must keep the prefix until +// the last reference is released (the reason the per-peer count must be an int, not a set). +func TestAllowedIPs_SamePeerMultipleRefs(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerA") + if f.adds != 1 { + t.Fatalf("expected a single wg add for the same peer, got %d", f.adds) + } + + mustDecrement(t, c, p, "peerA") + if f.installed[p] != "peerA" { + t.Fatalf("prefix must stay while a reference remains, got %q", f.installed[p]) + } + if f.removes != 0 { + t.Fatalf("no wg remove expected while a reference remains, got %d", f.removes) + } + + mustDecrement(t, c, p, "peerA") + if _, ok := f.installed[p]; ok { + t.Fatalf("prefix must be removed after last reference") + } +} + +// TestAllowedIPs_RefCountAndActive checks the Ref returned to callers (used for the HA-disabled log). +func TestAllowedIPs_RefCountAndActive(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + ref := mustIncrement(t, c, p, "peerA") + if ref.Count != 1 || ref.Out != "peerA" { + t.Fatalf("want {1, peerA}, got {%d, %q}", ref.Count, ref.Out) + } + ref = mustIncrement(t, c, p, "peerB") + if ref.Count != 2 || ref.Out != "peerA" { + t.Fatalf("want {2, peerA}, got {%d, %q}", ref.Count, ref.Out) + } +} + +// TestAllowedIPs_Flush removes everything installed and clears the counter. +func TestAllowedIPs_Flush(t *testing.T) { + f := newFakeWG() + c := f.counter() + p1 := mustPrefix(t, "10.44.8.0/24") + p2 := mustPrefix(t, "10.44.9.0/24") + + mustIncrement(t, c, p1, "peerA") + mustIncrement(t, c, p2, "peerB") + + if err := c.Flush(); err != nil { + t.Fatal(err) + } + if len(f.installed) != 0 { + t.Fatalf("expected all prefixes removed, got %v", f.installed) + } + // After flush, a fresh increment must add again. + mustIncrement(t, c, p1, "peerC") + if f.installed[p1] != "peerC" { + t.Fatalf("counter not reset after flush") + } +} + +// TestAllowedIPs_SelfHealAfterSwapAddError ensures a failed add during a swap does not permanently +// strand the prefix: the next Decrement (or Increment) must retry and install a surviving peer. +func TestAllowedIPs_SelfHealAfterSwapAddError(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + mustIncrement(t, c, p, "peerC") + + // Removing the active peerA triggers a swap to a survivor; make the add fail once. + f.failAdd = true + if _, err := c.Decrement(p, "peerA"); err == nil { + t.Fatalf("expected error from failed swap add") + } + if _, ok := f.installed[p]; ok { + t.Fatalf("nothing should be installed after a failed swap add, got %q", f.installed[p]) + } + + // A later Decrement of a non-active survivor must retry the hand-off (self-heal), not stay stuck. + ref := mustDecrement(t, c, p, "peerC") + if got := f.installed[p]; got == "" { + t.Fatalf("self-heal failed: prefix left unrouted after add recovered") + } + if ref.Out == "" { + t.Fatalf("expected an active peer after self-heal, got empty") + } +} + +// TestAllowedIPs_SelfHealAfterRemoveError ensures a failed remove during a swap is retried instead +// of leaving e.active stuck on a peer that no longer holds references. +func TestAllowedIPs_SelfHealAfterRemoveError(t *testing.T) { + f := newFakeWG() + c := f.counter() + p := mustPrefix(t, "10.44.8.0/24") + + mustIncrement(t, c, p, "peerA") + mustIncrement(t, c, p, "peerB") + + // Releasing active peerA must detach it (remove) then add peerB; fail the remove once. + f.failRemove = true + if _, err := c.Decrement(p, "peerA"); err == nil { + t.Fatalf("expected error from failed remove") + } + + // Next Decrement of the non-active survivor retries: removes stale peerA, installs peerB. + mustDecrement(t, c, p, "peerB") + // peerB had only one ref, so after retry the prefix is fully released. + if _, ok := f.installed[p]; ok { + t.Fatalf("expected prefix released after self-heal, still on %q", f.installed[p]) + } +} diff --git a/client/internal/routemanager/refcounter/types.go b/client/internal/routemanager/refcounter/types.go index aadac3e25..7da0e17e3 100644 --- a/client/internal/routemanager/refcounter/types.go +++ b/client/internal/routemanager/refcounter/types.go @@ -5,5 +5,7 @@ import "net/netip" // RouteRefCounter is a Counter for Route, it doesn't take any input on Increment and doesn't use any output on Decrement type RouteRefCounter = Counter[netip.Prefix, struct{}, struct{}] -// AllowedIPsRefCounter is a Counter for AllowedIPs, it takes a peer key on Increment and passes it back to Decrement -type AllowedIPsRefCounter = Counter[netip.Prefix, string, string] +// AllowedIPsRefCounter tracks WireGuard AllowedIPs per prefix. Unlike the generic Counter it is peer-aware: +// a prefix can be claimed by several peers at once and WireGuard allows a given prefix on exactly one peer, +// so the counter records the per-peer reference count and swaps the installed peer when the active one is released. +// See allowedips.go. diff --git a/client/internal/routemanager/static/route.go b/client/internal/routemanager/static/route.go index d480fdf00..8ba03d090 100644 --- a/client/internal/routemanager/static/route.go +++ b/client/internal/routemanager/static/route.go @@ -15,6 +15,11 @@ type Route struct { route *route.Route routeRefCounter *refcounter.RouteRefCounter allowedIPsRefcounter *refcounter.AllowedIPsRefCounter + // currentPeerKey is the routing peer this watcher currently has the prefix installed on + // (the HA winner elected by the watcher). It can differ from route.Peer and change on + // failover, so it is recorded on AddAllowedIPs and used on RemoveAllowedIPs to decrement + // the exact peer that was incremented. + currentPeerKey string } func NewRoute(params common.HandlerParams) *Route { @@ -52,12 +57,15 @@ func (r *Route) AddAllowedIPs(peerKey string) error { ref.Out, ) } + r.currentPeerKey = peerKey return nil } func (r *Route) RemoveAllowedIPs() error { - if _, err := r.allowedIPsRefcounter.Decrement(r.route.Network); err != nil { - return err + var err error + if _, decErr := r.allowedIPsRefcounter.Decrement(r.route.Network, r.currentPeerKey); decErr != nil { + err = fmt.Errorf("remove allowed IP %s: %w", r.route.Network, decErr) } - return nil + r.currentPeerKey = "" + return err } From 8a43f4f9439c239972644b4a6ec9440c929deb7e Mon Sep 17 00:00:00 2001 From: Riccardo Manfrin <3090891+riccardomanfrin@users.noreply.github.com> Date: Tue, 28 Jul 2026 09:42:57 +0200 Subject: [PATCH 08/32] [client] fix build: add ReapplyMatching to dedicated AllowedIPsRefCounter (#6935) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes #6799 turned `AllowedIPsRefCounter` from an alias of the generic `Counter` into a dedicated peer-aware type. A change merged in parallel — `DefaultManager.ReconcilePeerAllowedIPs` (lazy-connection idle→wake reconvergence) — calls `allowedIPsRefCounter.ReapplyMatching`, which only existed on the generic `Counter`. Each PR built alone; the merged `main` did not: ``` client/internal/routemanager/manager.go: m.allowedIPsRefCounter.ReapplyMatching undefined ``` Add `ReapplyMatching(pred, apply)` to the dedicated type, matching the generic contract (keyed on the active/`Out` peer): it re-applies every prefix whose currently installed peer satisfies `pred`, skipping prefixes with no active peer (reconciled by the next Increment/Decrement). Also update `reconcile_test.go` to construct the counter via `refcounter.NewAllowedIPs` — the generic `refcounter.New` return value is no longer assignable to the dedicated type. Verified with `cd client && CGO_ENABLED=1 go build .` (the failing CI step) and the `TestReconcilePeerAllowedIPs` / refcounter / routemanager tests. ## Issue ticket number and link No public issue — fixes a `main` build break from a semantic merge conflict between #6799 and the `ReconcilePeerAllowedIPs` change. Failing run: https://github.com/netbirdio/netbird/actions/runs/30330882389/job/90185591366 ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) Internal build fix restoring a method on the AllowedIPs refcounter. No public API, CLI, config, or behavior change. ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: N/A --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **Bug Fixes** * Improved route reconciliation so all applicable allowed IP prefixes are reliably re-applied for the correct peer. * Prevented reconciliation from affecting routes assigned to other peers. * Improved handling of errors encountered while restoring multiple routes, providing more consistent results. --- .../internal/routemanager/reconcile_test.go | 2 +- .../routemanager/refcounter/allowedips.go | 21 +++++++++++++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/client/internal/routemanager/reconcile_test.go b/client/internal/routemanager/reconcile_test.go index 2a8a4dc10..c6806a6cd 100644 --- a/client/internal/routemanager/reconcile_test.go +++ b/client/internal/routemanager/reconcile_test.go @@ -54,7 +54,7 @@ func (m *reconcileWGMock) GetNet() *netstack.Net { return n func TestReconcilePeerAllowedIPs(t *testing.T) { wg := &reconcileWGMock{} m := &DefaultManager{wgInterface: wg} - m.allowedIPsRefCounter = refcounter.New[netip.Prefix, string, string]( + m.allowedIPsRefCounter = refcounter.NewAllowedIPs( func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil }, func(netip.Prefix, string) error { return nil }, ) diff --git a/client/internal/routemanager/refcounter/allowedips.go b/client/internal/routemanager/refcounter/allowedips.go index 6f23162f6..6d682e8a9 100644 --- a/client/internal/routemanager/refcounter/allowedips.go +++ b/client/internal/routemanager/refcounter/allowedips.go @@ -169,6 +169,27 @@ func (rm *AllowedIPsRefCounter) Flush() error { return nberrors.FormatErrorOrNil(merr) } +// ReapplyMatching calls apply for every prefix whose currently installed (active) peer satisfies +// pred, holding the lock for the whole pass. It is used to re-push allowed IPs onto a peer whose +// WireGuard entry was rebuilt (e.g. a lazy connection cycling idle->wake) without a matching +// refcounter change, which would otherwise leave the prefix installed in the counter but missing +// on the device. Only the active peer is considered — a prefix that lost its installed peer to a +// failed swap is skipped here and reconciled by the next Increment/Decrement. +func (rm *AllowedIPsRefCounter) ReapplyMatching(pred func(out string) bool, apply func(key netip.Prefix) error) error { + rm.mu.Lock() + defer rm.mu.Unlock() + + var merr *multierror.Error + for prefix, e := range rm.entries { + if e.active != "" && pred(e.active) { + if err := apply(prefix); err != nil { + merr = multierror.Append(merr, err) + } + } + } + return nberrors.FormatErrorOrNil(merr) +} + // pickSurvivor deterministically selects a peer still referencing the prefix. WireGuard cannot do // multipath for a single prefix, so any surviving peer is a valid winner; the choice is made stable // (lowest peerKey) for predictable behavior and testability. From b3f9b82442a0a05b4c7ad737872c76ec5f43ee43 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Tue, 28 Jul 2026 21:45:57 +0900 Subject: [PATCH 09/32] [management] Force routing-peer DNS resolution for reverse-proxy domain targets (#6872) --- .../grpc/components_envelope_response.go | 2 +- .../internals/shared/grpc/conversion.go | 8 +-- .../internals/shared/grpc/conversion_test.go | 34 ++++++++++ management/internals/shared/grpc/server.go | 2 +- management/server/types/account.go | 48 +++++++++++++ management/server/types/account_components.go | 2 + management/server/types/account_test.go | 68 +++++++++++++++++++ shared/management/types/network.go | 5 ++ .../management/types/networkmap_components.go | 7 ++ 9 files changed, 170 insertions(+), 6 deletions(-) diff --git a/management/internals/shared/grpc/components_envelope_response.go b/management/internals/shared/grpc/components_envelope_response.go index cedd1b889..820708c98 100644 --- a/management/internals/shared/grpc/components_envelope_response.go +++ b/management/internals/shared/grpc/components_envelope_response.go @@ -50,7 +50,7 @@ func ToComponentSyncResponse( // TODO (dmitri) consider using invariants? // enableSSH := computeSSHEnabledForPeer(components, peer) - peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH) + peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH, components.ForceRoutingPeerDNSResolution) includeIPv6 := peer.SupportsIPv6() && peer.IPv6.IsValid() useSourcePrefixes := peer.SupportsSourcePrefixes() diff --git a/management/internals/shared/grpc/conversion.go b/management/internals/shared/grpc/conversion.go index 696d28f5c..74ceb3370 100644 --- a/management/internals/shared/grpc/conversion.go +++ b/management/internals/shared/grpc/conversion.go @@ -119,7 +119,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken return nbConfig } -func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool) *proto.PeerConfig { +func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig { netmask, _ := network.Net.Mask.Size() fqdn := peer.FQDN(dnsName) @@ -135,7 +135,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, set Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask), SshConfig: sshConfig, Fqdn: fqdn, - RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled, + RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled || peer.ProxyMeta.Embedded || forceRoutingPeerDNS, LazyConnectionEnabled: settings.LazyConnectionEnabled, AutoUpdate: &proto.AutoUpdateSettings{ Version: settings.AutoUpdateVersion, @@ -162,12 +162,12 @@ func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nb useSourcePrefixes := peer.SupportsSourcePrefixes() response := &proto.SyncResponse{ - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH), + PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution), NetworkMap: &proto.NetworkMap{ Serial: networkMap.Network.CurrentSerial(), Routes: networkmap.ToProtocolRoutes(networkMap.Routes), DNSConfig: networkmap.ToProtocolDNSConfig(networkMap.DNSConfig, dnsCache, dnsFwdPort), - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH), + PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution), }, Checks: toProtocolChecks(ctx, checks), } diff --git a/management/internals/shared/grpc/conversion_test.go b/management/internals/shared/grpc/conversion_test.go index 402b4fd07..38d370740 100644 --- a/management/internals/shared/grpc/conversion_test.go +++ b/management/internals/shared/grpc/conversion_test.go @@ -2,6 +2,7 @@ package grpc import ( "fmt" + "net" "net/netip" "reflect" "testing" @@ -14,6 +15,7 @@ import ( "github.com/netbirdio/netbird/management/internals/controllers/network_map" "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" + nbpeer "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/networkmap" ) @@ -301,3 +303,35 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) { assert.True(t, nbCfg.Metrics.Enabled, "metrics flag should carry the settings value") }) } + +func TestToPeerConfig_RoutingPeerDNSResolution(t *testing.T) { + network := &types.Network{Net: net.IPNet{IP: net.IPv4(100, 0, 0, 0), Mask: net.CIDRMask(8, 32)}} + + newPeer := func(embedded bool) *nbpeer.Peer { + p := &nbpeer.Peer{IP: netip.MustParseAddr("100.0.0.1")} + p.ProxyMeta.Embedded = embedded + return p + } + + tests := []struct { + name string + globalFlag bool + embedded bool + forceParam bool + wantEnabled bool + }{ + {name: "global off, regular peer, no force", wantEnabled: false}, + {name: "global on wins", globalFlag: true, wantEnabled: true}, + {name: "embedded proxy peer forced", embedded: true, wantEnabled: true}, + {name: "routing peer forced via param", forceParam: true, wantEnabled: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + settings := &types.Settings{RoutingPeerDNSResolutionEnabled: tt.globalFlag} + cfg := toPeerConfig(newPeer(tt.embedded), network, "netbird.selfhosted", settings, nil, nil, false, tt.forceParam) + assert.Equal(t, tt.wantEnabled, cfg.RoutingPeerDnsResolutionEnabled, + "RoutingPeerDnsResolutionEnabled should reflect global || embedded || forced") + }) + } +} diff --git a/management/internals/shared/grpc/server.go b/management/internals/shared/grpc/server.go index 3b7d62ac7..485f05a92 100644 --- a/management/internals/shared/grpc/server.go +++ b/management/internals/shared/grpc/server.go @@ -921,7 +921,7 @@ func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, ne // if peer has reached this point then it has logged in loginResp := &proto.LoginResponse{ NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, settings), - PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH), + PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false), Checks: toProtocolChecks(ctx, postureChecks), } diff --git a/management/server/types/account.go b/management/server/types/account.go index 588e63a09..1a3a30544 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -1517,6 +1517,54 @@ func (a *Account) GetResourceRoutersMap() map[string]map[string]*routerTypes.Net return routers } +// forcesRoutingPeerDNSResolution reports whether the given peer must run +// routing-peer DNS resolution regardless of the account-global +// RoutingPeerDNSResolutionEnabled setting. It returns true when the peer is a +// router for a domain network resource that is targeted by an enabled +// reverse-proxy service, so the peer's DNS forwarder starts and can resolve +// the target for the embedded proxy peers. Embedded proxy peers themselves are +// handled at PeerConfig build time. +func (a *Account) forcesRoutingPeerDNSResolution(peerID string, routers map[string]map[string]*routerTypes.NetworkRouter) bool { + targeted := a.proxyTargetedDomainResourceIDs() + if len(targeted) == 0 { + return false + } + + for _, resource := range a.NetworkResources { + if resource == nil || !resource.Enabled || resource.Type != resourceTypes.Domain { + continue + } + if _, ok := targeted[resource.ID]; !ok { + continue + } + if _, isRouter := routers[resource.NetworkID][peerID]; isRouter { + return true + } + } + + return false +} + +// proxyTargetedDomainResourceIDs returns the set of domain network resource IDs +// targeted by an enabled, non-terminated reverse-proxy service. +func (a *Account) proxyTargetedDomainResourceIDs() map[string]struct{} { + ids := make(map[string]struct{}) + for _, svc := range a.Services { + if svc == nil || !svc.Enabled || svc.Terminated { + continue + } + for _, target := range svc.Targets { + if target == nil || !target.Enabled { + continue + } + if target.TargetType == service.TargetTypeDomain { + ids[target.TargetId] = struct{}{} + } + } + } + return ids +} + // getPoliciesSourcePeers collects all unique peers from the source groups defined in the given policies. func getPoliciesSourcePeers(policies []*Policy, groups map[string]*Group) map[string]struct{} { sourcePeers := make(map[string]struct{}) diff --git a/management/server/types/account_components.go b/management/server/types/account_components.go index af27788d8..6fc904c0b 100644 --- a/management/server/types/account_components.go +++ b/management/server/types/account_components.go @@ -140,6 +140,8 @@ func (a *Account) GetPeerNetworkMapComponents( RouterPeers: make(map[string]*ComponentPeer), NetworkXIDToPublicID: make(map[string]string, len(a.Networks)), PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)), + + ForceRoutingPeerDNSResolution: a.forcesRoutingPeerDNSResolution(peerID, routers), } for _, n := range a.Networks { if n != nil { diff --git a/management/server/types/account_test.go b/management/server/types/account_test.go index 67d9e1c6f..80f2a950a 100644 --- a/management/server/types/account_test.go +++ b/management/server/types/account_test.go @@ -1751,3 +1751,71 @@ func hasPrivateAccessPolicy(account *Account, serviceID string) bool { } return false } + +func TestForcesRoutingPeerDNSResolution(t *testing.T) { + buildAccountRes := func(serviceEnabled, targetEnabled, resourceEnabled bool, targetType service.TargetType, resType resourceTypes.NetworkResourceType) *Account { + return &Account{ + Id: "accountID", + Groups: map[string]*Group{ + "router-group": {ID: "router-group", Peers: []string{"router-peer-grp"}}, + }, + NetworkRouters: []*routerTypes.NetworkRouter{ + {ID: "r1", NetworkID: "net-1", AccountID: "accountID", Peer: "router-peer", Enabled: true}, + {ID: "r2", NetworkID: "net-1", AccountID: "accountID", PeerGroups: []string{"router-group"}, Enabled: true}, + }, + NetworkResources: []*resourceTypes.NetworkResource{ + {ID: "res-domain", AccountID: "accountID", NetworkID: "net-1", Type: resType, Domain: "example.org", Enabled: resourceEnabled}, + }, + Services: []*service.Service{ + { + ID: "svc-1", AccountID: "accountID", Enabled: serviceEnabled, + Targets: []*service.Target{ + {TargetId: "res-domain", TargetType: targetType, Enabled: targetEnabled}, + }, + }, + }, + } + } + + buildAccount := func(serviceEnabled, targetEnabled, resourceEnabled bool, targetType service.TargetType) *Account { + return buildAccountRes(serviceEnabled, targetEnabled, resourceEnabled, targetType, resourceTypes.Domain) + } + + t.Run("router peer for RP-targeted domain resource is forced", func(t *testing.T) { + account := buildAccount(true, true, true, service.TargetTypeDomain) + routers := account.GetResourceRoutersMap() + assert.True(t, account.forcesRoutingPeerDNSResolution("router-peer", routers), "direct router peer should be forced") + assert.True(t, account.forcesRoutingPeerDNSResolution("router-peer-grp", routers), "group-member router peer should be forced") + }) + + t.Run("non-router peer is not forced", func(t *testing.T) { + account := buildAccount(true, true, true, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("other-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when service disabled", func(t *testing.T) { + account := buildAccount(false, true, true, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when target disabled", func(t *testing.T) { + account := buildAccount(true, false, true, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when resource disabled", func(t *testing.T) { + account := buildAccount(true, true, false, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced for non-domain target type", func(t *testing.T) { + account := buildAccount(true, true, true, service.TargetTypePeer) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when targeted resource is not a domain", func(t *testing.T) { + account := buildAccountRes(true, true, true, service.TargetTypeDomain, resourceTypes.Host) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()), + "a domain target pointing at a non-domain resource must not force resolution") + }) +} diff --git a/shared/management/types/network.go b/shared/management/types/network.go index 72a5cc5b3..34ce60436 100644 --- a/shared/management/types/network.go +++ b/shared/management/types/network.go @@ -47,6 +47,10 @@ type NetworkMap struct { ForwardingRules []*ForwardingRule AuthorizedUsers map[string]map[string]struct{} EnableSSH bool + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool } func (nm *NetworkMap) Merge(other *NetworkMap) { @@ -56,6 +60,7 @@ func (nm *NetworkMap) Merge(other *NetworkMap) { nm.FirewallRules = mergeUnique(nm.FirewallRules, other.FirewallRules) nm.RoutesFirewallRules = mergeUnique(nm.RoutesFirewallRules, other.RoutesFirewallRules) nm.ForwardingRules = mergeUnique(nm.ForwardingRules, other.ForwardingRules) + nm.ForceRoutingPeerDNSResolution = nm.ForceRoutingPeerDNSResolution || other.ForceRoutingPeerDNSResolution } type comparableObject[T any] interface { diff --git a/shared/management/types/networkmap_components.go b/shared/management/types/networkmap_components.go index a708e99e1..c4c437e4b 100644 --- a/shared/management/types/networkmap_components.go +++ b/shared/management/types/networkmap_components.go @@ -56,6 +56,11 @@ type NetworkMapComponents struct { // true when returning an empty-like map (returned instead of nil) empty bool + + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool } type routeIndexEntry struct { @@ -190,6 +195,8 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap { RoutesFirewallRules: append(networkResourcesFirewallRules, routesFirewallRules...), AuthorizedUsers: authorizedUsers, EnableSSH: sshEnabled, + + ForceRoutingPeerDNSResolution: c.ForceRoutingPeerDNSResolution, } } From 9269b56386234aaa1d58de67c864af63ec87de9e Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Tue, 28 Jul 2026 21:46:19 +0900 Subject: [PATCH 10/32] [management] Read reverse-proxy service and target columns in Postgres path (#6886) --- management/server/store/sql_store.go | 352 +++++++++++------- .../server/store/sql_store_pgx_parity_test.go | 74 ++++ .../server/store/sql_store_service_test.go | 89 +++++ 3 files changed, 380 insertions(+), 135 deletions(-) create mode 100644 management/server/store/sql_store_pgx_parity_test.go diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 3ad870ad3..670d9f781 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "fmt" + "math" "net" "net/netip" "net/url" @@ -2258,117 +2259,30 @@ func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*p return checks, nil } -func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { - const serviceQuery = `SELECT id, account_id, name, domain, enabled, auth, - meta_created_at, meta_certificate_issued_at, meta_status, proxy_cluster, - pass_host_header, rewrite_redirects, session_private_key, session_public_key, - mode, listen_port, port_auto_assigned, source, source_peer, terminated, - private, access_groups - FROM services WHERE account_id = $1` +// serviceSelectColumns and targetSelectColumns are the column lists the Postgres +// pgx read path scans. They must stay in sync with the rpservice.Service and +// rpservice.Target gorm models; TestPgxServiceColumnsMatchGorm enforces this. +const serviceSelectColumns = `id, account_id, name, domain, enabled, auth, restrictions, + meta_created_at, meta_certificate_issued_at, meta_last_renewed_at, meta_status, proxy_cluster, + pass_host_header, rewrite_redirects, session_private_key, session_public_key, + mode, listen_port, port_auto_assigned, source, source_peer, terminated, + private, access_groups` - const targetsQuery = `SELECT id, account_id, service_id, path, host, port, protocol, - target_id, target_type, enabled - FROM targets WHERE service_id = ANY($1)` +const targetSelectColumns = `id, account_id, service_id, path, host, port, protocol, + target_id, target_type, enabled, proxy_protocol, + skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers, + direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes, + capture_content_types, agent_network, disable_access_log` + +func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { + const serviceQuery = `SELECT ` + serviceSelectColumns + ` FROM services WHERE account_id = $1` serviceRows, err := s.pool.Query(ctx, serviceQuery, accountID) if err != nil { return nil, err } - services, err := pgx.CollectRows(serviceRows, func(row pgx.CollectableRow) (*rpservice.Service, error) { - var s rpservice.Service - var auth []byte - var accessGroups []byte - var createdAt, certIssuedAt sql.NullTime - var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString - var mode, source, sourcePeer sql.NullString - var terminated, portAutoAssigned, private sql.NullBool - var listenPort sql.NullInt64 - err := row.Scan( - &s.ID, - &s.AccountID, - &s.Name, - &s.Domain, - &s.Enabled, - &auth, - &createdAt, - &certIssuedAt, - &status, - &proxyCluster, - &s.PassHostHeader, - &s.RewriteRedirects, - &sessionPrivateKey, - &sessionPublicKey, - &mode, - &listenPort, - &portAutoAssigned, - &source, - &sourcePeer, - &terminated, - &private, - &accessGroups, - ) - if err != nil { - return nil, err - } - - if auth != nil { - if err := json.Unmarshal(auth, &s.Auth); err != nil { - return nil, err - } - } - - if len(accessGroups) > 0 { - if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { - return nil, fmt.Errorf("unmarshal access_groups: %w", err) - } - } - - if private.Valid { - s.Private = private.Bool - } - - s.Meta = rpservice.Meta{} - if createdAt.Valid { - s.Meta.CreatedAt = createdAt.Time - } - if certIssuedAt.Valid { - t := certIssuedAt.Time - s.Meta.CertificateIssuedAt = &t - } - if status.Valid { - s.Meta.Status = status.String - } - if proxyCluster.Valid { - s.ProxyCluster = proxyCluster.String - } - if sessionPrivateKey.Valid { - s.SessionPrivateKey = sessionPrivateKey.String - } - if sessionPublicKey.Valid { - s.SessionPublicKey = sessionPublicKey.String - } - if mode.Valid { - s.Mode = mode.String - } - if source.Valid { - s.Source = source.String - } - if sourcePeer.Valid { - s.SourcePeer = sourcePeer.String - } - if terminated.Valid { - s.Terminated = terminated.Bool - } - if portAutoAssigned.Valid { - s.PortAutoAssigned = portAutoAssigned.Bool - } - if listenPort.Valid { - s.ListenPort = uint16(listenPort.Int64) - } - s.Targets = []*rpservice.Target{} - return &s, nil - }) + services, err := pgx.CollectRows(serviceRows, scanService) if err != nil { return nil, err } @@ -2379,39 +2293,12 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv serviceIDs := make([]string, len(services)) serviceMap := make(map[string]*rpservice.Service) - for i, s := range services { - serviceIDs[i] = s.ID - serviceMap[s.ID] = s + for i, svc := range services { + serviceIDs[i] = svc.ID + serviceMap[svc.ID] = svc } - targetRows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) - if err != nil { - return nil, err - } - - targets, err := pgx.CollectRows(targetRows, func(row pgx.CollectableRow) (*rpservice.Target, error) { - var t rpservice.Target - var path sql.NullString - err := row.Scan( - &t.ID, - &t.AccountID, - &t.ServiceID, - &path, - &t.Host, - &t.Port, - &t.Protocol, - &t.TargetId, - &t.TargetType, - &t.Enabled, - ) - if err != nil { - return nil, err - } - if path.Valid { - t.Path = &path.String - } - return &t, nil - }) + targets, err := s.getServiceTargets(ctx, serviceIDs) if err != nil { return nil, err } @@ -2425,6 +2312,201 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv return services, nil } +func scanService(row pgx.CollectableRow) (*rpservice.Service, error) { + var s rpservice.Service + var auth []byte + var restrictions []byte + var accessGroups []byte + var createdAt, certIssuedAt, lastRenewedAt sql.NullTime + var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString + var mode, source, sourcePeer sql.NullString + var terminated, portAutoAssigned, private sql.NullBool + var listenPort sql.NullInt64 + err := row.Scan( + &s.ID, + &s.AccountID, + &s.Name, + &s.Domain, + &s.Enabled, + &auth, + &restrictions, + &createdAt, + &certIssuedAt, + &lastRenewedAt, + &status, + &proxyCluster, + &s.PassHostHeader, + &s.RewriteRedirects, + &sessionPrivateKey, + &sessionPublicKey, + &mode, + &listenPort, + &portAutoAssigned, + &source, + &sourcePeer, + &terminated, + &private, + &accessGroups, + ) + if err != nil { + return nil, err + } + + if auth != nil { + if err := json.Unmarshal(auth, &s.Auth); err != nil { + return nil, err + } + } + + if len(restrictions) > 0 { + if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil { + return nil, fmt.Errorf("unmarshal restrictions: %w", err) + } + } + + if len(accessGroups) > 0 { + if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { + return nil, fmt.Errorf("unmarshal access_groups: %w", err) + } + } + + if private.Valid { + s.Private = private.Bool + } + + s.Meta = serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt, status) + if proxyCluster.Valid { + s.ProxyCluster = proxyCluster.String + } + if sessionPrivateKey.Valid { + s.SessionPrivateKey = sessionPrivateKey.String + } + if sessionPublicKey.Valid { + s.SessionPublicKey = sessionPublicKey.String + } + if mode.Valid { + s.Mode = mode.String + } + if source.Valid { + s.Source = source.String + } + if sourcePeer.Valid { + s.SourcePeer = sourcePeer.String + } + if terminated.Valid { + s.Terminated = terminated.Bool + } + if portAutoAssigned.Valid { + s.PortAutoAssigned = portAutoAssigned.Bool + } + if listenPort.Valid { + if listenPort.Int64 < 0 || listenPort.Int64 > math.MaxUint16 { + return nil, fmt.Errorf("listen_port %d out of range", listenPort.Int64) + } + s.ListenPort = uint16(listenPort.Int64) + } + s.Targets = []*rpservice.Target{} + return &s, nil +} + +func serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt sql.NullTime, status sql.NullString) rpservice.Meta { + meta := rpservice.Meta{} + if createdAt.Valid { + meta.CreatedAt = createdAt.Time + } + if certIssuedAt.Valid { + t := certIssuedAt.Time + meta.CertificateIssuedAt = &t + } + if lastRenewedAt.Valid { + t := lastRenewedAt.Time + meta.LastRenewedAt = &t + } + if status.Valid { + meta.Status = status.String + } + return meta +} + +func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) { + const targetsQuery = `SELECT ` + targetSelectColumns + ` FROM targets WHERE service_id = ANY($1)` + + rows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) + if err != nil { + return nil, err + } + + return pgx.CollectRows(rows, scanTarget) +} + +func scanTarget(row pgx.CollectableRow) (*rpservice.Target, error) { + var t rpservice.Target + var path sql.NullString + var pathRewrite sql.NullString + var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool + var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64 + var customHeaders, middlewares, captureContentTypes []byte + err := row.Scan( + &t.ID, + &t.AccountID, + &t.ServiceID, + &path, + &t.Host, + &t.Port, + &t.Protocol, + &t.TargetId, + &t.TargetType, + &t.Enabled, + &proxyProtocol, + &skipTLSVerify, + &requestTimeout, + &sessionIdleTimeout, + &pathRewrite, + &customHeaders, + &directUpstream, + &middlewares, + &captureMaxRequestBytes, + &captureMaxResponseBytes, + &captureContentTypes, + &agentNetwork, + &disableAccessLog, + ) + if err != nil { + return nil, err + } + if path.Valid { + t.Path = &path.String + } + + t.ProxyProtocol = proxyProtocol.Bool + t.Options.SkipTLSVerify = skipTLSVerify.Bool + t.Options.RequestTimeout = time.Duration(requestTimeout.Int64) + t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64) + t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String) + t.Options.DirectUpstream = directUpstream.Bool + t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64 + t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64 + t.Options.AgentNetwork = agentNetwork.Bool + t.Options.DisableAccessLog = disableAccessLog.Bool + + if len(customHeaders) > 0 { + if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil { + return nil, fmt.Errorf("unmarshal custom_headers: %w", err) + } + } + if len(middlewares) > 0 { + if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil { + return nil, fmt.Errorf("unmarshal middlewares: %w", err) + } + } + if len(captureContentTypes) > 0 { + if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil { + return nil, fmt.Errorf("unmarshal capture_content_types: %w", err) + } + } + return &t, nil +} + func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) { const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1` rows, err := s.pool.Query(ctx, query, accountID) diff --git a/management/server/store/sql_store_pgx_parity_test.go b/management/server/store/sql_store_pgx_parity_test.go new file mode 100644 index 000000000..1f17817d0 --- /dev/null +++ b/management/server/store/sql_store_pgx_parity_test.go @@ -0,0 +1,74 @@ +package store + +import ( + "strings" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm/schema" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" +) + +// TestPgxServiceColumnsMatchGorm guards the Postgres pgx read path against +// drifting from the gorm model. The SQLite/MySQL gorm path loads rows by struct, +// so a new column on a model is picked up automatically, but the hand-written +// pgx SELECT in sql_store.go must be updated by hand. This test fails when a +// gorm column is missing from the pgx column list, which otherwise silently +// returns zero-valued on Postgres with no compile error. +func TestPgxServiceColumnsMatchGorm(t *testing.T) { + tests := []struct { + name string + model any + selectColumns string + // excluded lists gorm columns intentionally not loaded by the pgx path. + excluded map[string]struct{} + }{ + { + name: "service", + model: &rpservice.Service{}, + selectColumns: serviceSelectColumns, + }, + { + name: "target", + model: &rpservice.Target{}, + selectColumns: targetSelectColumns, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + selected := parseColumnList(tc.selectColumns) + for _, col := range gormColumnNames(t, tc.model) { + if _, ok := tc.excluded[col]; ok { + continue + } + _, ok := selected[col] + assert.Truef(t, ok, + "gorm column %q is not read by the Postgres pgx SELECT; add it to %sSelectColumns in sql_store.go (or to the test's excluded set if it is intentionally not loaded)", + col, tc.name) + } + }) + } +} + +func parseColumnList(cols string) map[string]struct{} { + set := make(map[string]struct{}) + for _, c := range strings.Split(cols, ",") { + if c = strings.TrimSpace(c); c != "" { + set[c] = struct{}{} + } + } + return set +} + +// gormColumnNames returns the DB column names gorm would migrate for the model, +// using the same default naming strategy the store configures. +func gormColumnNames(t *testing.T, model any) []string { + t.Helper() + sch, err := schema.Parse(model, &sync.Map{}, schema.NamingStrategy{}) + require.NoError(t, err) + return sch.DBNames +} diff --git a/management/server/store/sql_store_service_test.go b/management/server/store/sql_store_service_test.go index 34999da4b..0e14fbdab 100644 --- a/management/server/store/sql_store_service_test.go +++ b/management/server/store/sql_store_service_test.go @@ -5,6 +5,7 @@ import ( "os" "runtime" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -44,3 +45,91 @@ func TestSqlStore_GetAccount_PrivateServiceRoundtrip(t *testing.T) { assert.Equal(t, []string{"grp-admins", "grp-ops"}, got.AccessGroups) }) } + +// TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip guards the Postgres pgx +// read path (getServices) against silently dropping columns present on the gorm +// model. Before the fix these fields loaded correctly on SQLite but came back +// zero-valued on Postgres because the hand-written SELECT and scan omitted them. +func TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip(t *testing.T) { + if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skip("skip CI tests on darwin and windows") + } + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + account := newAccountWithId(ctx, "account_svc_opts", "testuser", "") + require.NoError(t, store.SaveAccount(ctx, account)) + + renewedAt := time.Now().UTC().Truncate(time.Second) + targetPath := "/api" + svc := &rpservice.Service{ + ID: "svc-opts", + AccountID: account.Id, + Name: "opts-svc", + Domain: "opts.example", + Enabled: true, + Mode: rpservice.ModeHTTP, + Restrictions: rpservice.AccessRestrictions{ + AllowedCIDRs: []string{"10.0.0.0/8"}, + BlockedCountries: []string{"XX"}, + CrowdSecMode: "block", + }, + Meta: rpservice.Meta{ + LastRenewedAt: &renewedAt, + }, + Targets: []*rpservice.Target{ + { + AccountID: account.Id, + ServiceID: "svc-opts", + Path: &targetPath, + Host: "backend.internal", + Port: 8080, + Protocol: "http", + TargetId: "tgt-1", + Enabled: true, + ProxyProtocol: true, + Options: rpservice.TargetOptions{ + SkipTLSVerify: true, + RequestTimeout: 30 * time.Second, + SessionIdleTimeout: 5 * time.Minute, + PathRewrite: rpservice.PathRewritePreserve, + CustomHeaders: map[string]string{"X-Foo": "bar"}, + DirectUpstream: true, + CaptureMaxRequestBytes: 1024, + CaptureMaxResponseBytes: 2048, + CaptureContentTypes: []string{"application/json"}, + AgentNetwork: true, + DisableAccessLog: true, + }, + }, + }, + } + require.NoError(t, store.CreateService(ctx, svc)) + + loaded, err := store.GetAccount(ctx, account.Id) + require.NoError(t, err) + require.Len(t, loaded.Services, 1) + + got := loaded.Services[0] + assert.Equal(t, []string{"10.0.0.0/8"}, got.Restrictions.AllowedCIDRs, "restrictions allowed CIDRs") + assert.Equal(t, []string{"XX"}, got.Restrictions.BlockedCountries, "restrictions blocked countries") + assert.Equal(t, "block", got.Restrictions.CrowdSecMode, "restrictions crowdsec mode") + require.NotNil(t, got.Meta.LastRenewedAt, "meta last renewed at") + assert.WithinDuration(t, renewedAt, *got.Meta.LastRenewedAt, time.Second, "meta last renewed at") + + require.Len(t, got.Targets, 1) + tg := got.Targets[0] + assert.True(t, tg.ProxyProtocol, "target proxy protocol") + assert.True(t, tg.Options.SkipTLSVerify, "options skip TLS verify") + assert.Equal(t, 30*time.Second, tg.Options.RequestTimeout, "options request timeout") + assert.Equal(t, 5*time.Minute, tg.Options.SessionIdleTimeout, "options session idle timeout") + assert.Equal(t, rpservice.PathRewritePreserve, tg.Options.PathRewrite, "options path rewrite") + assert.Equal(t, map[string]string{"X-Foo": "bar"}, tg.Options.CustomHeaders, "options custom headers") + assert.True(t, tg.Options.DirectUpstream, "options direct upstream") + assert.Equal(t, int64(1024), tg.Options.CaptureMaxRequestBytes, "options capture max request bytes") + assert.Equal(t, int64(2048), tg.Options.CaptureMaxResponseBytes, "options capture max response bytes") + assert.Equal(t, []string{"application/json"}, tg.Options.CaptureContentTypes, "options capture content types") + assert.True(t, tg.Options.AgentNetwork, "options agent network") + assert.True(t, tg.Options.DisableAccessLog, "options disable access log") + }) +} From 42e45ff9f95b4be38157f5c2d13bcd01349a4fb4 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 28 Jul 2026 16:22:48 +0200 Subject: [PATCH 11/32] [client] Expose RenameProfile in the Android profile manager binding (#6926) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Expose RenameProfile in the Android profile manager binding ## Issue ticket number and link ## Stack ### Checklist - [ ] Is it a bug fix - [ ] Is a typo/documentation fix - [x] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **New Features** * Added the ability to rename Android profiles. * Rename operations now provide clear success or failure feedback. --- client/android/profile_manager.go | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/client/android/profile_manager.go b/client/android/profile_manager.go index 87c001396..9a051137c 100644 --- a/client/android/profile_manager.go +++ b/client/android/profile_manager.go @@ -189,6 +189,19 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return nil } +// RenameProfile changes a profile's display name. The profile ID, and therefore +// its on-disk filename, is left untouched: only the "name" field of the config +// is rewritten. This works for the default profile too, whose config lives in +// netbird.cfg rather than under profiles/. +func (pm *ProfileManager) RenameProfile(id string, newName string) error { + if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), androidUsername, newName); err != nil { + return fmt.Errorf("failed to rename profile: %w", err) + } + + log.Infof("renamed profile %s to: %s", id, newName) + return nil +} + // RemoveProfile deletes a profile func (pm *ProfileManager) RemoveProfile(id string) error { // Use ServiceManager (removes profile from profiles/ directory) From 0fb4c8c42330e1cb6698d9e270876e34f9e62e8f Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 28 Jul 2026 16:23:09 +0200 Subject: [PATCH 12/32] [client] Build UI release binaries with the production tag (#6898) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes The goreleaser UI configs passed no build tags, so released netbird-ui binaries were built on the Wails !production path: DevTools enabled, browser context menu forced on, and FRONTEND_DEVSERVER_URL still honored. ## Issue ticket number and link ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **Chores** * Updated release build configurations to mark Linux, Windows, and macOS UI builds as production releases. --- .goreleaser_ui.yaml | 6 ++++++ .goreleaser_ui_darwin.yaml | 2 ++ 2 files changed, 8 insertions(+) diff --git a/.goreleaser_ui.yaml b/.goreleaser_ui.yaml index 197fcd440..1157e6379 100644 --- a/.goreleaser_ui.yaml +++ b/.goreleaser_ui.yaml @@ -24,6 +24,8 @@ builds: ldflags: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production - id: netbird-ui-windows-amd64 dir: client/ui @@ -39,6 +41,8 @@ builds: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser - -H windowsgui mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production - id: netbird-ui-windows-arm64 dir: client/ui @@ -55,6 +59,8 @@ builds: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser - -H windowsgui mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production archives: - id: linux-arch diff --git a/.goreleaser_ui_darwin.yaml b/.goreleaser_ui_darwin.yaml index 96e15371a..47b991344 100644 --- a/.goreleaser_ui_darwin.yaml +++ b/.goreleaser_ui_darwin.yaml @@ -29,6 +29,8 @@ builds: ldflags: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production universal_binaries: - id: netbird-ui-darwin From 4acbe2670af3d7881c85698f662285066be34304 Mon Sep 17 00:00:00 2001 From: Stefan Fast Date: Tue, 28 Jul 2026 16:24:30 +0200 Subject: [PATCH 13/32] [client] Escape dots in interface names for sysctl configuration (#6930) --- .../internal/routemanager/sysctl/sysctl_linux.go | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/client/internal/routemanager/sysctl/sysctl_linux.go b/client/internal/routemanager/sysctl/sysctl_linux.go index f96a57f37..46b7c9fb7 100644 --- a/client/internal/routemanager/sysctl/sysctl_linux.go +++ b/client/internal/routemanager/sysctl/sysctl_linux.go @@ -20,6 +20,8 @@ const ( rpFilterPath = "net.ipv4.conf.all.rp_filter" rpFilterInterfacePath = "net.ipv4.conf.%s.rp_filter" srcValidMarkPath = "net.ipv4.conf.all.src_valid_mark" + percentEscape = "%25" + dotEscape = "%2E" ) type iface interface { @@ -56,7 +58,11 @@ func Setup(wgIface iface) (map[string]int, error) { continue } - i := fmt.Sprintf(rpFilterInterfacePath, intf.Name) + // Escape '%' and '.' so they survive the dot-to-slash conversion in Set() + safeName := strings.ReplaceAll(intf.Name, "%", percentEscape) + safeName = strings.ReplaceAll(safeName, ".", dotEscape) + + i := fmt.Sprintf(rpFilterInterfacePath, safeName) oldVal, err := Set(i, 2, true) if err != nil { result = multierror.Append(result, err) @@ -70,7 +76,11 @@ func Setup(wgIface iface) (map[string]int, error) { // Set sets a sysctl configuration, if onlyIfOne is true it will only set the new value if it's set to 1 func Set(key string, desiredValue int, onlyIfOne bool) (int, error) { - path := fmt.Sprintf("/proc/sys/%s", strings.ReplaceAll(key, ".", "/")) + path := strings.ReplaceAll(key, ".", "/") + // Unescape interface dots and percent signs + path = strings.ReplaceAll(path, dotEscape, ".") + path = strings.ReplaceAll(path, percentEscape, "%") + path = fmt.Sprintf("/proc/sys/%s", path) currentValue, err := os.ReadFile(path) if err != nil { return -1, fmt.Errorf("read sysctl %s: %w", key, err) From 63c320b6a9c62e734f715a546dd06ba2de2afe62 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 28 Jul 2026 16:51:01 +0200 Subject: [PATCH 14/32] [client] Serialize iOS tunnel reconfiguration callbacks (#6870) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit On iOS the tunnel reconfiguration (setTunnelNetworkSettings) was driven from three independent Go paths without serialization: route prefix updates via the route notifier's own delivery loop, interface IP/IPv6 set synchronously from the engine start goroutine, and DNS config applied from the DNS apply chain through a separate Swift object. The three sources mutated the shared Swift settings-manager state concurrently and triggered overlapping updateTunnel() calls, losing or half-applying route and DNS settings. Introduce client/internal/tunnelnotifier: a single notifier that implements both listener.NetworkChangeListener and dns.IosDnsManager, queues all four callbacks (OnNetworkChanged, SetInterfaceIP, SetInterfaceIPv6, ApplyDns) in one FIFO and delivers them one-by-one from a single goroutine, so calls into Swift never overlap and arrive in order. RunOniOS wraps the two Swift objects into the notifier and closes it after the run loop exits. The route notifier keeps its prefix dedup but delegates delivery to the shared notifier instead of maintaining its own queue and loop; the dedup check and the enqueue stay under one mutex so queue order matches state-update order. Setting the interface IP becomes asynchronous, which is safe: on iOS wgInterface.Create() only uses the TunFd, and the FIFO preserves the IP -> routes/DNS relative order. The package is build-tag free so the unit tests run on any platform under -race. Android is unaffected: it receives all settings atomically in one configureInterface call and serializes TUN rebuilds on a single handler thread. ## Describe your changes ## Issue ticket number and link ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [x] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **New Features** * Improved iOS tunnel networking by coordinating network, interface, route, and DNS updates through a unified notifier. * Updated iOS behavior so route/prefix changes are applied immediately instead of via queued/background delivery. * **Bug Fixes** * Improved reliability by ensuring network and DNS-related callbacks are invoked in order without overlap. * Ensured pending updates are drained and handled correctly during shutdown. * **Tests** * Added coverage for FIFO ordering, non-overlapping callback execution, interleaved DNS/route updates, and graceful shutdown behavior. --- client/internal/connect.go | 8 +- client/internal/mobile_dependency.go | 12 +- .../routemanager/notifier/notifier_ios.go | 44 +--- client/internal/tunnelnotifier/notifier.go | 124 +++++++++++ .../internal/tunnelnotifier/notifier_test.go | 192 ++++++++++++++++++ 5 files changed, 334 insertions(+), 46 deletions(-) create mode 100644 client/internal/tunnelnotifier/notifier.go create mode 100644 client/internal/tunnelnotifier/notifier_test.go diff --git a/client/internal/connect.go b/client/internal/connect.go index f4d14aab2..87126b222 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -34,6 +34,7 @@ import ( "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/statemanager" "github.com/netbirdio/netbird/client/internal/stdnet" + "github.com/netbirdio/netbird/client/internal/tunnelnotifier" "github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/internal/updater/installer" nbnet "github.com/netbirdio/netbird/client/net" @@ -136,10 +137,13 @@ func (c *ConnectClient) RunOniOS( // Set GC percent to 5% to reduce memory usage as iOS only allows 50MB of memory for the extension. debug.SetGCPercent(5) + notifier := tunnelnotifier.New(networkChangeListener, dnsManager) + defer notifier.Close() + mobileDependency := MobileDependency{ FileDescriptor: fileDescriptor, - NetworkChangeListener: networkChangeListener, - DnsManager: dnsManager, + NetworkChangeListener: notifier, + DnsManager: notifier, StateFilePath: stateFilePath, TempDir: cacheDir, } diff --git a/client/internal/mobile_dependency.go b/client/internal/mobile_dependency.go index 310d61a25..0234432b1 100644 --- a/client/internal/mobile_dependency.go +++ b/client/internal/mobile_dependency.go @@ -11,12 +11,14 @@ import ( // MobileDependency collect all dependencies for mobile platform type MobileDependency struct { - // Android only - TunAdapter device.TunAdapter - IFaceDiscover stdnet.ExternalIFaceDiscover + // Android and iOS NetworkChangeListener listener.NetworkChangeListener - HostDNSAddresses []netip.AddrPort - DnsReadyListener dns.ReadyListener + + // Android only + TunAdapter device.TunAdapter + IFaceDiscover stdnet.ExternalIFaceDiscover + HostDNSAddresses []netip.AddrPort + DnsReadyListener dns.ReadyListener // iOS only DnsManager dns.IosDnsManager diff --git a/client/internal/routemanager/notifier/notifier_ios.go b/client/internal/routemanager/notifier/notifier_ios.go index d0888f3a1..c91a76551 100644 --- a/client/internal/routemanager/notifier/notifier_ios.go +++ b/client/internal/routemanager/notifier/notifier_ios.go @@ -3,7 +3,6 @@ package notifier import ( - "container/list" "net/netip" "slices" "sort" @@ -16,20 +15,12 @@ import ( type Notifier struct { mu sync.Mutex - cond *sync.Cond currentPrefixes []string listener listener.NetworkChangeListener - queue *list.List - closed bool } func NewNotifier() *Notifier { - n := &Notifier{ - queue: list.New(), - } - n.cond = sync.NewCond(&n.mu) - go n.deliverLoop() - return n + return &Notifier{} } func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { @@ -59,44 +50,19 @@ func (n *Notifier) OnNewPrefixes(prefixes []netip.Prefix) { sort.Strings(newNets) n.mu.Lock() + defer n.mu.Unlock() if slices.Equal(n.currentPrefixes, newNets) { - n.mu.Unlock() return } n.currentPrefixes = newNets - routes := strings.Join(n.currentPrefixes, ",") - n.queue.PushBack(routes) - n.cond.Signal() - n.mu.Unlock() + if n.listener != nil { + n.listener.OnNetworkChanged(strings.Join(n.currentPrefixes, ",")) + } } func (n *Notifier) Close() { - n.mu.Lock() - n.closed = true - n.cond.Signal() - n.mu.Unlock() } func (n *Notifier) GetInitialRouteRanges() []string { return nil } - -func (n *Notifier) deliverLoop() { - for { - n.mu.Lock() - for n.queue.Len() == 0 && !n.closed { - n.cond.Wait() - } - if n.closed && n.queue.Len() == 0 { - n.mu.Unlock() - return - } - routes := n.queue.Remove(n.queue.Front()).(string) - l := n.listener - n.mu.Unlock() - - if l != nil { - l.OnNetworkChanged(routes) - } - } -} diff --git a/client/internal/tunnelnotifier/notifier.go b/client/internal/tunnelnotifier/notifier.go new file mode 100644 index 000000000..b62923a6e --- /dev/null +++ b/client/internal/tunnelnotifier/notifier.go @@ -0,0 +1,124 @@ +package tunnelnotifier + +import ( + "container/list" + "sync" + + "github.com/netbirdio/netbird/client/internal/dns" + "github.com/netbirdio/netbird/client/internal/listener" +) + +type eventKind int + +const ( + eventRoutes eventKind = iota + eventIfaceIP + eventIfaceIPv6 + eventDNS +) + +var ( + _ listener.NetworkChangeListener = (*Notifier)(nil) + _ dns.IosDnsManager = (*Notifier)(nil) +) + +type event struct { + kind eventKind + payload string +} + +type Notifier struct { + mu sync.Mutex + cond *sync.Cond + queue *list.List + closed bool + done chan struct{} + + listener listener.NetworkChangeListener + dnsManager dns.IosDnsManager +} + +func New(l listener.NetworkChangeListener, dm dns.IosDnsManager) *Notifier { + n := &Notifier{ + queue: list.New(), + done: make(chan struct{}), + listener: l, + dnsManager: dm, + } + n.cond = sync.NewCond(&n.mu) + go n.deliverLoop() + return n +} + +func (n *Notifier) OnNetworkChanged(routes string) { + n.enqueue(event{kind: eventRoutes, payload: routes}) +} + +func (n *Notifier) SetInterfaceIP(ip string) { + n.enqueue(event{kind: eventIfaceIP, payload: ip}) +} + +func (n *Notifier) SetInterfaceIPv6(ip string) { + n.enqueue(event{kind: eventIfaceIPv6, payload: ip}) +} + +func (n *Notifier) ApplyDns(config string) { + n.enqueue(event{kind: eventDNS, payload: config}) +} + +// Close stops accepting new events and blocks until the delivery loop has +// drained all queued events and exited. +func (n *Notifier) Close() { + n.mu.Lock() + n.closed = true + n.cond.Signal() + n.mu.Unlock() + <-n.done +} + +func (n *Notifier) enqueue(ev event) { + n.mu.Lock() + defer n.mu.Unlock() + if n.closed { + return + } + n.queue.PushBack(ev) + n.cond.Signal() +} + +func (n *Notifier) deliverLoop() { + defer close(n.done) + for { + n.mu.Lock() + for n.queue.Len() == 0 && !n.closed { + n.cond.Wait() + } + if n.closed && n.queue.Len() == 0 { + n.mu.Unlock() + return + } + ev := n.queue.Remove(n.queue.Front()).(event) + l := n.listener + dm := n.dnsManager + n.mu.Unlock() + + switch ev.kind { + case eventRoutes: + if l != nil { + l.OnNetworkChanged(ev.payload) + } + case eventIfaceIP: + if l != nil { + l.SetInterfaceIP(ev.payload) + } + case eventIfaceIPv6: + if l != nil { + l.SetInterfaceIPv6(ev.payload) + } + case eventDNS: + if dm != nil { + dm.ApplyDns(ev.payload) + } + } + } +} diff --git a/client/internal/tunnelnotifier/notifier_test.go b/client/internal/tunnelnotifier/notifier_test.go new file mode 100644 index 000000000..ffbcdc15c --- /dev/null +++ b/client/internal/tunnelnotifier/notifier_test.go @@ -0,0 +1,192 @@ +package tunnelnotifier + +import ( + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type call struct { + kind string + payload string +} + +type recorder struct { + mu sync.Mutex + calls []call + inFlight atomic.Int32 + overlap atomic.Bool + delay time.Duration +} + +func (r *recorder) record(kind, payload string) { + if r.inFlight.Add(1) != 1 { + r.overlap.Store(true) + } + if r.delay > 0 { + time.Sleep(r.delay) + } + r.mu.Lock() + r.calls = append(r.calls, call{kind: kind, payload: payload}) + r.mu.Unlock() + r.inFlight.Add(-1) +} + +func (r *recorder) count() int { + r.mu.Lock() + defer r.mu.Unlock() + return len(r.calls) +} + +func (r *recorder) snapshot() []call { + r.mu.Lock() + defer r.mu.Unlock() + out := make([]call, len(r.calls)) + copy(out, r.calls) + return out +} + +type fakeListener struct { + rec *recorder +} + +func (f *fakeListener) OnNetworkChanged(routes string) { + f.rec.record("routes", routes) +} + +func (f *fakeListener) SetInterfaceIP(ip string) { + f.rec.record("ip", ip) +} + +func (f *fakeListener) SetInterfaceIPv6(ip string) { + f.rec.record("ipv6", ip) +} + +type fakeDNSManager struct { + rec *recorder +} + +func (f *fakeDNSManager) ApplyDns(config string) { + f.rec.record("dns", config) +} + +func TestFIFOOrder(t *testing.T) { + rec := &recorder{} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + defer n.Close() + + n.SetInterfaceIP("10.0.0.1") + n.SetInterfaceIPv6("fd00::1") + n.ApplyDns(`{"domains":[]}`) + n.OnNetworkChanged("10.0.0.0/8,192.168.0.0/16") + n.ApplyDns(`{"domains":["example.com"]}`) + + require.Eventually(t, func() bool { return rec.count() == 5 }, time.Second, time.Millisecond) + + expected := []call{ + {kind: "ip", payload: "10.0.0.1"}, + {kind: "ipv6", payload: "fd00::1"}, + {kind: "dns", payload: `{"domains":[]}`}, + {kind: "routes", payload: "10.0.0.0/8,192.168.0.0/16"}, + {kind: "dns", payload: `{"domains":["example.com"]}`}, + } + assert.Equal(t, expected, rec.snapshot()) +} + +func TestNoOverlappingCalls(t *testing.T) { + rec := &recorder{delay: 100 * time.Microsecond} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + defer n.Close() + + const producers = 8 + const perProducer = 25 + + var wg sync.WaitGroup + for i := 0; i < producers; i++ { + wg.Add(1) + go func(id int) { + defer wg.Done() + for j := 0; j < perProducer; j++ { + payload := fmt.Sprintf("%d-%d", id, j) + switch j % 4 { + case 0: + n.OnNetworkChanged(payload) + case 1: + n.SetInterfaceIP(payload) + case 2: + n.SetInterfaceIPv6(payload) + case 3: + n.ApplyDns(payload) + } + } + }(i) + } + wg.Wait() + + require.Eventually(t, func() bool { return rec.count() == producers*perProducer }, 5*time.Second, time.Millisecond) + assert.False(t, rec.overlap.Load()) +} + +func TestDNSAndRoutesInterleaved(t *testing.T) { + rec := &recorder{delay: 100 * time.Microsecond} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + defer n.Close() + + const events = 50 + + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + for i := 0; i < events; i++ { + n.ApplyDns(fmt.Sprintf("dns-%d", i)) + } + }() + go func() { + defer wg.Done() + for i := 0; i < events; i++ { + n.OnNetworkChanged(fmt.Sprintf("routes-%d", i)) + } + }() + wg.Wait() + + require.Eventually(t, func() bool { return rec.count() == 2*events }, 5*time.Second, time.Millisecond) + assert.False(t, rec.overlap.Load()) + + var dnsSeen, routesSeen int + for _, c := range rec.snapshot() { + switch c.kind { + case "dns": + assert.Equal(t, fmt.Sprintf("dns-%d", dnsSeen), c.payload) + dnsSeen++ + case "routes": + assert.Equal(t, fmt.Sprintf("routes-%d", routesSeen), c.payload) + routesSeen++ + } + } + assert.Equal(t, events, dnsSeen) + assert.Equal(t, events, routesSeen) +} + +func TestCloseDrainsQueue(t *testing.T) { + rec := &recorder{delay: time.Millisecond} + n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec}) + + const events = 20 + for i := 0; i < events; i++ { + n.OnNetworkChanged(fmt.Sprintf("routes-%d", i)) + } + n.Close() + + require.Equal(t, events, rec.count(), "Close must not return before all queued events are delivered") + + n.OnNetworkChanged("after-close") + n.ApplyDns("after-close") + time.Sleep(50 * time.Millisecond) + assert.Equal(t, events, rec.count()) +} From 2ef457be957cd24e4b329b894635b389c62eb27c Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 28 Jul 2026 17:39:37 +0200 Subject: [PATCH 15/32] [client] Unify route selection in the route manager (#6928) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Move route select/deselect handling from the daemon server into exported routemanager methods (SelectRoutes, DeselectRoutes, SelectAllRoutes, DeselectAllRoutes) so every consumer shares one implementation: v4/v6 exit-pair expansion, exit-node mutual exclusion, and selection triggering. Previously the exit-node exclusivity lived only in the daemon's SelectNetworks RPC, so the Android and iOS bindings could leave two exit nodes selected until the next network map reconciliation. Both bindings now call the shared manager methods and enforce exclusivity at toggle time, matching the desktop behavior. ## Issue ticket number and link ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **New Features** * Route selection/deselection is now handled through shared route-manager APIs for both individual routes and “all routes”. * Exit-node selections automatically enforce mutual exclusivity while keeping non-exit routes unaffected. * **Bug Fixes** * Unknown or unavailable route IDs now return errors, and exclusivity is preserved even when some route IDs fail. * **Tests** * Added route-selection tests covering exclusivity (including IPv4/IPv6), select-all behavior, partial errors, and invalid IDs. * **Refactor / Chores** * Simplified Android, iOS, and server routing flows to delegate to the shared manager; updated mocks and removed redundant routing command logic/dependencies. --- client/android/client.go | 12 +- client/android/route_command.go | 70 --------- client/internal/routemanager/manager.go | 13 +- client/internal/routemanager/mock.go | 26 ++++ client/internal/routemanager/selection.go | 138 ++++++++++++++++++ .../internal/routemanager/selection_test.go | 129 ++++++++++++++++ client/ios/NetBirdSDK/client.go | 42 ++---- client/server/network.go | 74 +--------- client/server/network_exitnode_test.go | 26 ---- 9 files changed, 324 insertions(+), 206 deletions(-) delete mode 100644 client/android/route_command.go create mode 100644 client/internal/routemanager/selection.go create mode 100644 client/internal/routemanager/selection_test.go delete mode 100644 client/server/network_exitnode_test.go diff --git a/client/android/client.go b/client/android/client.go index 2266ff53d..2ce627f10 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -439,10 +439,6 @@ func (c *Client) RemoveConnectionListener() { c.recorder.RemoveConnectionListener() } -func (c *Client) toggleRoute(command routeCommand) error { - return command.toggleRoute() -} - func (c *Client) getRouteManager() (routemanager.Manager, error) { client := c.getConnectClient() if client == nil { @@ -462,22 +458,22 @@ func (c *Client) getRouteManager() (routemanager.Manager, error) { return manager, nil } -func (c *Client) SelectRoute(route string) error { +func (c *Client) SelectRoute(id string) error { manager, err := c.getRouteManager() if err != nil { return err } - return c.toggleRoute(selectRouteCommand{route: route, manager: manager}) + return manager.SelectRoutes([]route.NetID{route.NetID(id)}, true) } -func (c *Client) DeselectRoute(route string) error { +func (c *Client) DeselectRoute(id string) error { manager, err := c.getRouteManager() if err != nil { return err } - return c.toggleRoute(deselectRouteCommand{route: route, manager: manager}) + return manager.DeselectRoutes([]route.NetID{route.NetID(id)}) } // getNetworkDomainsFromRoute extracts domains from a route and enriches each domain diff --git a/client/android/route_command.go b/client/android/route_command.go deleted file mode 100644 index 5e7357335..000000000 --- a/client/android/route_command.go +++ /dev/null @@ -1,70 +0,0 @@ -//go:build android - -package android - -import ( - "fmt" - - log "github.com/sirupsen/logrus" - "golang.org/x/exp/maps" - - "github.com/netbirdio/netbird/client/internal/routemanager" - "github.com/netbirdio/netbird/route" -) - -func executeRouteToggle(id string, manager routemanager.Manager, - operationName string, - routeOperation func(routes []route.NetID, allRoutes []route.NetID) error) error { - netID := route.NetID(id) - routes := []route.NetID{netID} - - routesMap := manager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - - log.Debugf("%s with ids: %v", operationName, routes) - - if err := routeOperation(routes, maps.Keys(routesMap)); err != nil { - log.Debugf("error when %s: %s", operationName, err) - return fmt.Errorf("error %s: %w", operationName, err) - } - - manager.TriggerSelection(manager.GetClientRoutes()) - - return nil -} - -type routeCommand interface { - toggleRoute() error -} - -type selectRouteCommand struct { - route string - manager routemanager.Manager -} - -func (s selectRouteCommand) toggleRoute() error { - routeSelector := s.manager.GetRouteSelector() - if routeSelector == nil { - return fmt.Errorf("no route selector available") - } - - routeOperation := func(routes []route.NetID, allRoutes []route.NetID) error { - return routeSelector.SelectRoutes(routes, true, allRoutes) - } - - return executeRouteToggle(s.route, s.manager, "selecting route", routeOperation) -} - -type deselectRouteCommand struct { - route string - manager routemanager.Manager -} - -func (d deselectRouteCommand) toggleRoute() error { - routeSelector := d.manager.GetRouteSelector() - if routeSelector == nil { - return fmt.Errorf("no route selector available") - } - - return executeRouteToggle(d.route, d.manager, "deselecting route", routeSelector.DeselectRoutes) -} diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go index 7a818b539..2ab7e2a85 100644 --- a/client/internal/routemanager/manager.go +++ b/client/internal/routemanager/manager.go @@ -52,6 +52,10 @@ type Manager interface { UpdateRoutes(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error ClassifyRoutes(newRoutes []*route.Route) (map[route.ID]*route.Route, route.HAMap) TriggerSelection(route.HAMap) + SelectRoutes(ids []route.NetID, appendRoute bool) error + DeselectRoutes(ids []route.NetID) error + SelectAllRoutes() + DeselectAllRoutes() GetRouteSelector() *routeselector.RouteSelector GetClientRoutes() route.HAMap GetSelectedClientRoutes() route.HAMap @@ -800,7 +804,7 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI var info exitNodeInfo for haID, routes := range clientRoutes { - if !m.isExitNodeRoute(routes) { + if !isExitNodeRoutes(routes) { continue } @@ -820,13 +824,6 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI return info } -func (m *DefaultManager) isExitNodeRoute(routes []*route.Route) bool { - if len(routes) == 0 { - return false - } - return route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network) -} - func (m *DefaultManager) categorizeUserSelection(netID route.NetID, info *exitNodeInfo) { if m.routeSelector.IsSelected(netID) { info.userSelected = append(info.userSelected, netID) diff --git a/client/internal/routemanager/mock.go b/client/internal/routemanager/mock.go index c1620b24c..cf761091d 100644 --- a/client/internal/routemanager/mock.go +++ b/client/internal/routemanager/mock.go @@ -16,6 +16,8 @@ type MockManager struct { ClassifyRoutesFunc func(routes []*route.Route) (map[route.ID]*route.Route, route.HAMap) UpdateRoutesFunc func(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error TriggerSelectionFunc func(haMap route.HAMap) + SelectRoutesFunc func(ids []route.NetID, appendRoute bool) error + DeselectRoutesFunc func(ids []route.NetID) error GetRouteSelectorFunc func() *routeselector.RouteSelector GetClientRoutesFunc func() route.HAMap GetSelectedClientRoutesFunc func() route.HAMap @@ -55,6 +57,30 @@ func (m *MockManager) TriggerSelection(networks route.HAMap) { } } +// SelectRoutes mock implementation of SelectRoutes from Manager interface +func (m *MockManager) SelectRoutes(ids []route.NetID, appendRoute bool) error { + if m.SelectRoutesFunc != nil { + return m.SelectRoutesFunc(ids, appendRoute) + } + return nil +} + +// DeselectRoutes mock implementation of DeselectRoutes from Manager interface +func (m *MockManager) DeselectRoutes(ids []route.NetID) error { + if m.DeselectRoutesFunc != nil { + return m.DeselectRoutesFunc(ids) + } + return nil +} + +// SelectAllRoutes mock implementation of SelectAllRoutes from Manager interface +func (m *MockManager) SelectAllRoutes() { +} + +// DeselectAllRoutes mock implementation of DeselectAllRoutes from Manager interface +func (m *MockManager) DeselectAllRoutes() { +} + // GetRouteSelector mock implementation of GetRouteSelector from Manager interface func (m *MockManager) GetRouteSelector() *routeselector.RouteSelector { if m.GetRouteSelectorFunc != nil { diff --git a/client/internal/routemanager/selection.go b/client/internal/routemanager/selection.go new file mode 100644 index 000000000..6d5feec79 --- /dev/null +++ b/client/internal/routemanager/selection.go @@ -0,0 +1,138 @@ +package routemanager + +import ( + "fmt" + "slices" + + "github.com/hashicorp/go-multierror" + log "github.com/sirupsen/logrus" + "golang.org/x/exp/maps" + + nberrors "github.com/netbirdio/netbird/client/errors" + "github.com/netbirdio/netbird/route" +) + +// SelectRoutes selects the routes with the given network IDs and applies the +// new selection. V4/v6 exit-node pairs are expanded automatically. Exit nodes +// are mutually exclusive: if the selection activates an exit node, every other +// available exit node is deselected so two can't be active at once. With +// appendRoute=false the previous selection is replaced instead of extended. +func (m *DefaultManager) SelectRoutes(ids []route.NetID, appendRoute bool) error { + if err := m.selectRoutes(ids, appendRoute); err != nil { + return err + } + m.TriggerSelection(m.GetClientRoutes()) + return nil +} + +// DeselectRoutes removes the routes with the given network IDs from the +// selection and applies the change. V4/v6 exit-node pairs are expanded +// automatically. +func (m *DefaultManager) DeselectRoutes(ids []route.NetID) error { + if err := m.deselectRoutes(ids); err != nil { + return err + } + m.TriggerSelection(m.GetClientRoutes()) + return nil +} + +func (m *DefaultManager) deselectRoutes(ids []route.NetID) error { + routesMap := m.GetClientRoutesWithNetID() + routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap) + + log.Debugf("deselecting routes with ids: %v", routes) + + if err := m.routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil { + return fmt.Errorf("deselect routes: %w", err) + } + + return nil +} + +// SelectAllRoutes selects every available route and applies the selection. +// Exit nodes stay mutually exclusive: at most one remains active. +func (m *DefaultManager) SelectAllRoutes() { + m.selectAllRoutes() + m.TriggerSelection(m.GetClientRoutes()) +} + +func (m *DefaultManager) selectAllRoutes() { + m.routeSelector.SelectAllRoutes() + + // Select-all wipes every explicit selection, so exit nodes fall back to + // management's auto-apply flags — which may mark several at once. + // Reconcile immediately so at most one exit node stays active instead of + // waiting for the next network map to enforce it. + m.mux.Lock() + defer m.mux.Unlock() + m.updateRouteSelectorFromManagement(m.clientRoutes) +} + +// DeselectAllRoutes deselects every route and applies the change. +func (m *DefaultManager) DeselectAllRoutes() { + m.routeSelector.DeselectAllRoutes() + m.TriggerSelection(m.GetClientRoutes()) +} + +func (m *DefaultManager) selectRoutes(ids []route.NetID, appendRoute bool) error { + routesMap := m.GetClientRoutesWithNetID() + routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap) + allIDs := maps.Keys(routesMap) + + log.Debugf("selecting routes with ids: %v", routes) + + // A partial failure (e.g. an unknown ID in the request) still selects the + // valid routes, so exclusivity below must run regardless of the error. + var merr *multierror.Error + if err := m.routeSelector.SelectRoutes(routes, appendRoute, allIDs); err != nil { + merr = multierror.Append(merr, fmt.Errorf("select routes: %w", err)) + } + + // Exit nodes are mutually exclusive: if this selection activates an + // exit node, deselect every other available exit node so two can't be + // selected at once. Non-exit route selections are left untouched. + if requestActivatesExitNode(routes, routesMap) { + if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 { + if err := m.routeSelector.DeselectRoutes(others, allIDs); err != nil { + merr = multierror.Append(merr, fmt.Errorf("deselect sibling exit nodes: %w", err)) + } + } + } + + return nberrors.FormatErrorOrNil(merr) +} + +func isExitNodeRoutes(routes []*route.Route) bool { + return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network)) +} + +// requestActivatesExitNode reports whether any requested NetID maps to an exit +// node (default route) in the current route table. +func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool { + for _, id := range requested { + if isExitNodeRoutes(routesMap[id]) { + return true + } + } + return false +} + +// otherExitNodeIDs returns every available exit-node NetID that is not in the +// requested set — the siblings to deselect so a single exit node stays active. +func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID { + keep := make(map[route.NetID]struct{}, len(requested)) + for _, id := range requested { + keep[id] = struct{}{} + } + var others []route.NetID + for id, routes := range routesMap { + if !isExitNodeRoutes(routes) { + continue + } + if _, ok := keep[id]; ok { + continue + } + others = append(others, id) + } + return others +} diff --git a/client/internal/routemanager/selection_test.go b/client/internal/routemanager/selection_test.go new file mode 100644 index 000000000..6066b5661 --- /dev/null +++ b/client/internal/routemanager/selection_test.go @@ -0,0 +1,129 @@ +package routemanager + +import ( + "net/netip" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/routeselector" + "github.com/netbirdio/netbird/route" +) + +func v6ExitRoute(netID, peer string) *route.Route { + return &route.Route{ + NetID: route.NetID(netID), + Network: netip.MustParsePrefix("::/0"), + Peer: peer, + } +} + +func newSelectionTestManager() *DefaultManager { + return &DefaultManager{ + routeSelector: routeselector.NewRouteSelector(), + clientRoutes: route.HAMap{ + "exitA|0.0.0.0/0": {exitRoute("exitA", "p1", true)}, + "exitA-v6|::/0": {v6ExitRoute("exitA-v6", "p1")}, + "exitB|0.0.0.0/0": {exitRoute("exitB", "p2", true)}, + "lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}}, + }, + } +} + +func TestSelectRoutes_ExitNodeExclusivity(t *testing.T) { + m := newSelectionTestManager() + + // Selecting an exit node selects its v6 pair and deselects the sibling. + require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true)) + assert.True(t, m.routeSelector.IsSelected("exitA"), "exitA should be selected") + assert.True(t, m.routeSelector.IsSelected("exitA-v6"), "the v6 pair follows its v4 base") + assert.False(t, m.routeSelector.IsSelected("exitB"), "the sibling exit node must be deselected") + + // Switching to the sibling deselects the previous exit node and its v6 pair. + require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true)) + assert.True(t, m.routeSelector.IsSelected("exitB"), "exitB should now be selected") + assert.False(t, m.routeSelector.IsSelected("exitA"), "the previous exit node must be deselected") + assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "the previous exit node's v6 pair must be deselected") + assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched") + + // Selecting a non-exit route leaves the active exit node alone. + require.NoError(t, m.selectRoutes([]route.NetID{"lan"}, true)) + assert.True(t, m.routeSelector.IsSelected("exitB"), "selecting a non-exit route keeps the exit node") + + // Deselecting the active exit node turns every exit node off. + require.NoError(t, m.deselectRoutes([]route.NetID{"exitB"})) + assert.False(t, m.routeSelector.IsSelected("exitB"), "exitB should be deselected") + assert.False(t, m.routeSelector.IsSelected("exitA"), "exitA stays deselected") + assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched") +} + +func TestSelectRoutes_PartialErrorStillEnforcesExclusivity(t *testing.T) { + // The unknown ID must be reported, but the valid exit node in the same + // request is still selected — so its sibling must still be deselected. + // Both orderings are covered: processing must continue past the invalid + // ID wherever it sits in the request. + requests := map[string][]route.NetID{ + "invalid id first": {"missing", "exitB"}, + "invalid id last": {"exitB", "missing"}, + } + + for name, ids := range requests { + t.Run(name, func(t *testing.T) { + m := newSelectionTestManager() + + require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true)) + + err := m.selectRoutes(ids, true) + assert.Error(t, err, "unknown id must be reported") + assert.True(t, m.routeSelector.IsSelected("exitB"), "valid exit node from the request is selected") + assert.False(t, m.routeSelector.IsSelected("exitA"), "sibling exit node must be deselected despite the error") + assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "sibling's v6 pair must be deselected too") + }) + } +} + +func TestSelectAllRoutes_KeepsSingleExitNode(t *testing.T) { + // Both exit nodes are marked for auto-apply by management + // (SkipAutoApply=false), the state where select-all could turn on two at + // once without the immediate reconciliation. + m := &DefaultManager{ + routeSelector: routeselector.NewRouteSelector(), + clientRoutes: route.HAMap{ + "exitA|0.0.0.0/0": {exitRoute("exitA", "p1", false)}, + "exitB|0.0.0.0/0": {exitRoute("exitB", "p2", false)}, + "lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}}, + }, + } + + require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true)) + + m.selectAllRoutes() + + assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit routes are all selected") + assert.True(t, m.routeSelector.IsSelected("exitA"), "the deterministic management pick stays active") + assert.False(t, m.routeSelector.IsSelected("exitB"), "select-all must not leave a second exit node active") +} + +func TestSelectRoutes_UnknownRoute(t *testing.T) { + m := newSelectionTestManager() + + assert.Error(t, m.selectRoutes([]route.NetID{"missing"}, true), "selecting an unavailable route must fail") + assert.Error(t, m.deselectRoutes([]route.NetID{"missing"}), "deselecting an unavailable route must fail") +} + +func TestExitNodeSelectionHelpers(t *testing.T) { + routesMap := map[route.NetID][]*route.Route{ + "exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}}, + "exitB": {{Network: netip.MustParsePrefix("::/0")}}, + "lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}}, + } + + assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node") + assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node") + assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node") + assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node") + + others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"}) + assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored") +} diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index a2f123900..9289a3910 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -13,7 +13,6 @@ import ( "time" log "github.com/sirupsen/logrus" - "golang.org/x/exp/maps" "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/auth" @@ -637,23 +636,18 @@ func (c *Client) SelectRoute(id string) error { } routeManager := engine.GetRouteManager() - routeSelector := routeManager.GetRouteSelector() if id == "All" { log.Debugf("select all routes") - routeSelector.SelectAllRoutes() - } else { - log.Debugf("select route with id: %s", id) - routes := toNetIDs([]string{id}) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - if err := routeSelector.SelectRoutes(routes, true, maps.Keys(routesMap)); err != nil { - log.Debugf("error when selecting routes: %s", err) - return fmt.Errorf("select routes: %w", err) - } + routeManager.SelectAllRoutes() + return nil } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) - return nil + log.Debugf("select route with id: %s", id) + if err := routeManager.SelectRoutes(toNetIDs([]string{id}), true); err != nil { + log.Debugf("error when selecting routes: %s", err) + return err + } + return nil } func (c *Client) DeselectRoute(id string) error { @@ -667,21 +661,17 @@ func (c *Client) DeselectRoute(id string) error { } routeManager := engine.GetRouteManager() - routeSelector := routeManager.GetRouteSelector() if id == "All" { log.Debugf("deselect all routes") - routeSelector.DeselectAllRoutes() - } else { - log.Debugf("deselect route with id: %s", id) - routes := toNetIDs([]string{id}) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - if err := routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil { - log.Debugf("error when deselecting routes: %s", err) - return fmt.Errorf("deselect routes: %w", err) - } + routeManager.DeselectAllRoutes() + return nil + } + + log.Debugf("deselect route with id: %s", id) + if err := routeManager.DeselectRoutes(toNetIDs([]string{id})); err != nil { + log.Debugf("error when deselecting routes: %s", err) + return err } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) return nil } diff --git a/client/server/network.go b/client/server/network.go index c38715256..c390b8180 100644 --- a/client/server/network.go +++ b/client/server/network.go @@ -8,7 +8,6 @@ import ( "sort" "strings" - "golang.org/x/exp/maps" "google.golang.org/grpc/codes" gstatus "google.golang.org/grpc/status" @@ -161,30 +160,11 @@ func (s *Server) SelectNetworks(_ context.Context, req *proto.SelectNetworksRequ return nil, fmt.Errorf("no route manager") } - routeSelector := routeManager.GetRouteSelector() if req.GetAll() { - routeSelector.SelectAllRoutes() - } else { - routes := toNetIDs(req.GetNetworkIDs()) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - netIdRoutes := maps.Keys(routesMap) - if err := routeSelector.SelectRoutes(routes, req.GetAppend(), netIdRoutes); err != nil { - return nil, fmt.Errorf("select routes: %w", err) - } - - // Exit nodes are mutually exclusive: if this selection activates an - // exit node, deselect every other available exit node so two can't be - // selected at once. Non-exit route selections are left untouched. - if requestActivatesExitNode(routes, routesMap) { - if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 { - if err := routeSelector.DeselectRoutes(others, netIdRoutes); err != nil { - return nil, fmt.Errorf("deselect sibling exit nodes: %w", err) - } - } - } + routeManager.SelectAllRoutes() + } else if err := routeManager.SelectRoutes(toNetIDs(req.GetNetworkIDs()), req.GetAppend()); err != nil { + return nil, err } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) s.statusRecorder.PublishEvent( proto.SystemEvent_INFO, @@ -224,19 +204,11 @@ func (s *Server) DeselectNetworks(_ context.Context, req *proto.SelectNetworksRe return nil, fmt.Errorf("no route manager") } - routeSelector := routeManager.GetRouteSelector() if req.GetAll() { - routeSelector.DeselectAllRoutes() - } else { - routes := toNetIDs(req.GetNetworkIDs()) - routesMap := routeManager.GetClientRoutesWithNetID() - routes = route.ExpandV6ExitPairs(routes, routesMap) - netIdRoutes := maps.Keys(routesMap) - if err := routeSelector.DeselectRoutes(routes, netIdRoutes); err != nil { - return nil, fmt.Errorf("deselect routes: %w", err) - } + routeManager.DeselectAllRoutes() + } else if err := routeManager.DeselectRoutes(toNetIDs(req.GetNetworkIDs())); err != nil { + return nil, err } - routeManager.TriggerSelection(routeManager.GetClientRoutes()) s.statusRecorder.PublishEvent( proto.SystemEvent_INFO, @@ -261,37 +233,3 @@ func toNetIDs(routes []string) []route.NetID { return netIDs } -func isExitNodeRoutes(routes []*route.Route) bool { - return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network)) -} - -// requestActivatesExitNode reports whether any requested NetID maps to an exit -// node (default route) in the current route table. -func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool { - for _, id := range requested { - if isExitNodeRoutes(routesMap[id]) { - return true - } - } - return false -} - -// otherExitNodeIDs returns every available exit-node NetID that is not in the -// requested set — the siblings to deselect so a single exit node stays active. -func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID { - keep := make(map[route.NetID]struct{}, len(requested)) - for _, id := range requested { - keep[id] = struct{}{} - } - var others []route.NetID - for id, routes := range routesMap { - if !isExitNodeRoutes(routes) { - continue - } - if _, ok := keep[id]; ok { - continue - } - others = append(others, id) - } - return others -} diff --git a/client/server/network_exitnode_test.go b/client/server/network_exitnode_test.go deleted file mode 100644 index 1c0ba0ecb..000000000 --- a/client/server/network_exitnode_test.go +++ /dev/null @@ -1,26 +0,0 @@ -package server - -import ( - "net/netip" - "testing" - - "github.com/stretchr/testify/assert" - - "github.com/netbirdio/netbird/route" -) - -func TestExitNodeSelectionHelpers(t *testing.T) { - routesMap := map[route.NetID][]*route.Route{ - "exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}}, - "exitB": {{Network: netip.MustParsePrefix("::/0")}}, - "lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}}, - } - - assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node") - assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node") - assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node") - assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node") - - others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"}) - assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored") -} From 3d1f209ea38a80f11e3141651148cda3c3f45e13 Mon Sep 17 00:00:00 2001 From: Riccardo Manfrin <3090891+riccardomanfrin@users.noreply.github.com> Date: Tue, 28 Jul 2026 18:12:49 +0200 Subject: [PATCH 16/32] [client] parse NB_LAZY_CONN_INACTIVITY_THRESHOLD as a Go duration (#6947) ## Describe your changes `NB_LAZY_CONN_INACTIVITY_THRESHOLD` was parsed with `strconv.Atoi`, i.e. as a bare integer number of minutes. The documentation, however, states it takes a Go duration (e.g. `30m`, `1h`). As a result any documented value such as `30m` or `5m` failed to parse and **silently fell back to the 15m default**, so the setting appeared to have no effect. `inactivityThresholdEnv()` now parses the value with `time.ParseDuration`, matching the docs. A bare integer is still accepted as a number of minutes for backwards compatibility, and an unparseable value logs a warning and falls back to the default. Added `TestInactivityThresholdEnv` covering Go-duration values (`30m`/`1h`/`90s`), the bare-integer minutes fallback, and zero/negative/garbage inputs. ## Issue ticket number and link N/A --- client/internal/conn_mgr.go | 19 ++++++++++++----- client/internal/conn_mgr_test.go | 35 ++++++++++++++++++++++++++++++++ 2 files changed, 49 insertions(+), 5 deletions(-) diff --git a/client/internal/conn_mgr.go b/client/internal/conn_mgr.go index 7a591f60c..ad0f00c5d 100644 --- a/client/internal/conn_mgr.go +++ b/client/internal/conn_mgr.go @@ -385,11 +385,20 @@ func inactivityThresholdEnv() *time.Duration { return nil } - parsedMinutes, err := strconv.Atoi(envValue) - if err != nil || parsedMinutes <= 0 { - return nil + // Documented format: a Go duration such as "30m" or "1h". + if d, err := time.ParseDuration(envValue); err == nil { + if d <= 0 { + return nil + } + return &d } - d := time.Duration(parsedMinutes) * time.Minute - return &d + // Backwards compatibility: a bare integer used to be interpreted as minutes. + if parsedMinutes, err := strconv.Atoi(envValue); err == nil && parsedMinutes > 0 { + d := time.Duration(parsedMinutes) * time.Minute + return &d + } + + log.Warnf("invalid %s value %q: expected a Go duration such as 30m or 1h", lazyconn.EnvInactivityThreshold, envValue) + return nil } diff --git a/client/internal/conn_mgr_test.go b/client/internal/conn_mgr_test.go index e027fd4f2..ac5d6f2c8 100644 --- a/client/internal/conn_mgr_test.go +++ b/client/internal/conn_mgr_test.go @@ -104,3 +104,38 @@ func TestConnMgr_ActivatePeerConcurrentWithLifecycle(t *testing.T) { close(done) wg.Wait() } + +func TestInactivityThresholdEnv(t *testing.T) { + tests := []struct { + name string + val string + want *time.Duration + }{ + {name: "unset", val: "", want: nil}, + {name: "go duration minutes", val: "30m", want: durPtr(30 * time.Minute)}, + {name: "go duration hours", val: "1h", want: durPtr(time.Hour)}, + {name: "go duration seconds", val: "90s", want: durPtr(90 * time.Second)}, + {name: "bare integer is minutes (backwards compat)", val: "5", want: durPtr(5 * time.Minute)}, + {name: "zero duration", val: "0s", want: nil}, + {name: "zero integer", val: "0", want: nil}, + {name: "negative duration", val: "-5m", want: nil}, + {name: "garbage", val: "abc", want: nil}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Setenv(lazyconn.EnvInactivityThreshold, tc.val) + got := inactivityThresholdEnv() + switch { + case tc.want == nil && got != nil: + t.Fatalf("want nil, got %v", *got) + case tc.want != nil && got == nil: + t.Fatalf("want %v, got nil", *tc.want) + case tc.want != nil && *got != *tc.want: + t.Fatalf("want %v, got %v", *tc.want, *got) + } + }) + } +} + +func durPtr(d time.Duration) *time.Duration { return &d } From 44fef45c2f48e3da14d3d1b9ee49ea6227f442b7 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 28 Jul 2026 18:13:22 +0200 Subject: [PATCH 17/32] [client] Export peer details for Android (#6925) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Eexport peer details for Android ## Issue ticket number and link ## Stack ### Checklist - [x] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **New Features** * Enhanced Android peer details with latency, data transfer totals, connection timestamps, handshake information, relay and security status, and ICE connection details. * Added formatted latency values to support clearer connection-quality indicators. * **Bug Fixes** * Peer lists now refresh WireGuard statistics before displaying transfer and handshake information, ensuring current values are shown. --- client/android/client.go | 47 +++++++++++++++++++++++++++++++++ client/android/peer_notifier.go | 18 +++++++++++++ 2 files changed, 65 insertions(+) diff --git a/client/android/client.go b/client/android/client.go index 2ce627f10..b3a845818 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -7,6 +7,7 @@ import ( "fmt" "os" "slices" + "strings" "sync" "time" @@ -299,6 +300,13 @@ func (c *Client) SetInfoLogLevel() { // PeersList return with the list of the PeerInfos func (c *Client) PeersList() *PeerInfoArray { + // The recorder only caches transfer counters and handshake times; nothing + // refreshes them on its own, so without this they read as zero. The desktop + // daemon does the same before serving a full peer status. + if err := c.recorder.RefreshWireGuardStats(); err != nil { + log.Debugf("failed to refresh WireGuard stats: %v", err) + } + fullStatus := c.recorder.GetFullStatus() peerInfos := make([]PeerInfo, len(fullStatus.Peers)) @@ -309,6 +317,20 @@ func (c *Client) PeersList() *PeerInfoArray { FQDN: p.FQDN, ConnStatus: int(p.ConnStatus), Routes: PeerRoutes{routes: maps.Keys(p.GetRoutes())}, + + PubKey: p.PubKey, + Latency: formatDuration(p.Latency), + LatencyMs: p.Latency.Milliseconds(), + BytesRx: p.BytesRx, + BytesTx: p.BytesTx, + ConnStatusUpdate: formatTime(p.ConnStatusUpdate), + Relayed: p.Relayed, + RosenpassEnabled: p.RosenpassEnabled, + LastWireguardHandshake: formatTime(p.LastWireguardHandshake), + LocalIceCandidateType: p.LocalIceCandidateType, + RemoteIceCandidateType: p.RemoteIceCandidateType, + LocalIceCandidateEndpoint: p.LocalIceCandidateEndpoint, + RemoteIceCandidateEndpoint: p.RemoteIceCandidateEndpoint, } peerInfos[n] = pi } @@ -508,3 +530,28 @@ func exportEnvList(list *EnvList) { } } } + +// formatDuration renders a duration for display, trimming the fractional part +// to two digits so latencies read as "12.34ms" rather than "12.345678ms". +func formatDuration(d time.Duration) string { + ds := d.String() + dotIndex := strings.Index(ds, ".") + if dotIndex == -1 { + return ds + } + + endIndex := min(dotIndex+3, len(ds)) + + // Skip the remaining digits so only the unit suffix is appended back. + unitStart := endIndex + for unitStart < len(ds) && ds[unitStart] >= '0' && ds[unitStart] <= '9' { + unitStart++ + } + return ds[:endIndex] + ds[unitStart:] +} + +// formatTime renders a timestamp in UTC using a fixed layout. The zero time is +// passed through as-is so the UI can recognise it and show "never" instead. +func formatTime(t time.Time) string { + return t.UTC().Format("2006-01-02 15:04:05") +} diff --git a/client/android/peer_notifier.go b/client/android/peer_notifier.go index c2595e574..f525055bb 100644 --- a/client/android/peer_notifier.go +++ b/client/android/peer_notifier.go @@ -12,12 +12,30 @@ const ( ) // PeerInfo describe information about the peers. It designed for the UI usage +// +// The fields below ConnStatus back the peer detail screen. Durations and times +// are pre-formatted into strings so the UI does not have to know Go's layouts; +// Latency is additionally exposed as LatencyMs for colour coding. type PeerInfo struct { IP string IPv6 string FQDN string ConnStatus int Routes PeerRoutes + + PubKey string + Latency string + LatencyMs int64 + BytesRx int64 + BytesTx int64 + ConnStatusUpdate string + Relayed bool + RosenpassEnabled bool + LastWireguardHandshake string + LocalIceCandidateType string + RemoteIceCandidateType string + LocalIceCandidateEndpoint string + RemoteIceCandidateEndpoint string } func (p *PeerInfo) GetPeerRoutes() *PeerRoutes { From dd2bdc0de3aa14dd14dc90611c69c420bb3d3eb2 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Tue, 28 Jul 2026 18:34:55 +0200 Subject: [PATCH 18/32] [management] explicit accountID check when deleting a user (#6944) --- management/server/user.go | 4 +++ management/server/user_test.go | 46 ++++++++++++++++++++++++++++++++++ 2 files changed, 50 insertions(+) diff --git a/management/server/user.go b/management/server/user.go index 1de63c302..fc8400e29 100644 --- a/management/server/user.go +++ b/management/server/user.go @@ -321,6 +321,10 @@ func (am *DefaultAccountManager) DeleteUser(ctx context.Context, accountID, init return err } + if targetUser.AccountID != accountID { + return status.NewUserNotFoundError(targetUserID) + } + if targetUser.Role == types.UserRoleOwner { return status.NewOwnerDeletePermissionError() } diff --git a/management/server/user_test.go b/management/server/user_test.go index f32a6b3a1..a2e71616a 100644 --- a/management/server/user_test.go +++ b/management/server/user_test.go @@ -802,6 +802,52 @@ func TestUser_DeleteUser_SelfDelete(t *testing.T) { } } +func TestUser_DeleteUser_OtherAccount(t *testing.T) { + testStore, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + if err != nil { + t.Fatalf("Error when creating store: %s", err) + } + t.Cleanup(cleanup) + + account := newAccountWithId(context.Background(), mockAccountID, mockUserID, "", "", "", false) + if err = testStore.SaveAccount(context.Background(), account); err != nil { + t.Fatalf("Error when saving account: %s", err) + } + + otherAccount := newAccountWithId(context.Background(), "otherAccount", "otherOwner", "", "", "", false) + otherAccount.Users["otherRegularUser"] = &types.User{ + Id: "otherRegularUser", + AccountID: "otherAccount", + Role: types.UserRoleUser, + } + otherAccount.Users["otherServiceUser"] = &types.User{ + Id: "otherServiceUser", + AccountID: "otherAccount", + Role: types.UserRoleUser, + IsServiceUser: true, + ServiceUserName: "otherServiceUser", + } + if err = testStore.SaveAccount(context.Background(), otherAccount); err != nil { + t.Fatalf("Error when saving other account: %s", err) + } + + am := DefaultAccountManager{ + Store: testStore, + eventStore: &activity.InMemoryEventStore{}, + permissionsManager: permissions.NewManager(testStore), + } + + for _, targetUserID := range []string{"otherRegularUser", "otherServiceUser"} { + t.Run(targetUserID, func(t *testing.T) { + err := am.DeleteUser(context.Background(), mockAccountID, mockUserID, targetUserID) + assert.Equal(t, status.NewUserNotFoundError(targetUserID), err) + + _, err = testStore.GetUserByUserID(context.Background(), store.LockingStrengthNone, targetUserID) + assert.NoError(t, err, "user of another account must not be deleted") + }) + } +} + func TestUser_DeleteUser_regularUser(t *testing.T) { store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) if err != nil { From df39c2b254ca7443afbfb8ff30638433466db230 Mon Sep 17 00:00:00 2001 From: Brandon Hopkins <76761586+TechHutTV@users.noreply.github.com> Date: Wed, 29 Jul 2026 03:20:43 -0700 Subject: [PATCH 19/32] [infrastructure] Deprecate legacy Dex and Zitadel getting-started scripts (#6952) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Retire the legacy Dex and Zitadel installation scripts by replacing their implementations with compatibility notices that: - Exit unsuccessfully before making any system changes. - Direct new deployments to `getting-started.sh` and the self-hosting quickstart. - Explain that the current installer uses NetBird's embedded Dex-based identity provider. - Direct users to configure Zitadel through the NetBird Dashboard or follow the advanced guide for a standalone identity-provider deployment. - Clarify that Dex support, Zitadel support, and existing deployments are not deprecated. - Announce that the compatibility notices will be removed in NetBird v0.80. Replace the legacy Zitadel provisioning workflow with CI checks that verify both scripts fail, write their notices to stderr, and include the expected documentation links and removal version. ## Issue ticket number and link No single issue tracks this deprecation. Related to: - Legacy installer failures: #4980, #4698, #3566, #3010, #2834, #2094, #2046, and #1966. - Legacy installer enhancement requests: #4510 and #1785. ## Stack Standalone PR based on `main`. ### Checklist - [ ] Is it a bug fix - [ ] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [x] Created tests that fail without the change (if possible) - [x] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change because the compatibility notices link to the existing quickstart, identity-provider, and advanced deployment documentation. ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: N/A --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit - **Breaking Changes** - Retired the legacy Dex and Zitadel getting-started installation scripts. - These scripts now exit with an explanatory deprecation message and direct users to the current setup flow and documentation. - Existing workflows now verify that the retired scripts fail as expected and provide the appropriate migration guidance. --- .../workflows/test-infrastructure-files.yml | 87 +- .../getting-started-with-dex.sh | 560 +--------- .../getting-started-with-zitadel.sh | 997 +----------------- 3 files changed, 48 insertions(+), 1596 deletions(-) diff --git a/.github/workflows/test-infrastructure-files.yml b/.github/workflows/test-infrastructure-files.yml index 0a4f2e371..965d8aa5d 100644 --- a/.github/workflows/test-infrastructure-files.yml +++ b/.github/workflows/test-infrastructure-files.yml @@ -249,78 +249,35 @@ jobs: docker compose exec management ls -l /var/lib/netbird/ | grep -i GeoLite2-City_[0-9]*.mmdb docker compose exec management ls -l /var/lib/netbird/ | grep -i geonames_[0-9]*.db - test-getting-started-script: + test-legacy-getting-started-scripts: runs-on: ubuntu-latest steps: - - name: Install jq - run: sudo apt-get install -y jq - - name: Checkout code uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: persist-credentials: false - - name: run script with Zitadel PostgreSQL - run: NETBIRD_DOMAIN=use-ip bash -x infrastructure_files/getting-started-with-zitadel.sh - - - name: test Caddy file gen postgres - run: test -f Caddyfile - - - name: test docker-compose file gen postgres - run: test -f docker-compose.yml - - - name: test management.json file gen postgres - run: test -f management.json - - - name: test turnserver.conf file gen postgres + - name: Verify Dex retirement notice run: | - set -x - test -f turnserver.conf - grep external-ip turnserver.conf + if infrastructure_files/getting-started-with-dex.sh >stdout.txt 2>stderr.txt; then + echo "Expected the retired Dex installer to fail" + exit 1 + fi + test ! -s stdout.txt + grep -Fq "Dex support is not deprecated." stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-quickstart" stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/identity-providers/local" stderr.txt + grep -Fq "removed in NetBird v0.80" stderr.txt - - name: test zitadel.env file gen postgres - run: test -f zitadel.env - - - name: test dashboard.env file gen postgres - run: test -f dashboard.env - - - name: test relay.env file gen postgres - run: test -f relay.env - - - name: test zdb.env file gen postgres - run: test -f zdb.env - - - name: Postgres run cleanup + - name: Verify Zitadel retirement notice run: | - docker compose down --volumes --rmi all - rm -rf docker-compose.yml Caddyfile zitadel.env dashboard.env machinekey/zitadel-admin-sa.token turnserver.conf management.json zdb.env - - - name: run script with Zitadel CockroachDB - run: bash -x infrastructure_files/getting-started-with-zitadel.sh - env: - NETBIRD_DOMAIN: use-ip - ZITADEL_DATABASE: cockroach - - - name: test Caddy file gen CockroachDB - run: test -f Caddyfile - - - name: test docker-compose file gen CockroachDB - run: test -f docker-compose.yml - - - name: test management.json file gen CockroachDB - run: test -f management.json - - - name: test turnserver.conf file gen CockroachDB - run: | - set -x - test -f turnserver.conf - grep external-ip turnserver.conf - - - name: test zitadel.env file gen CockroachDB - run: test -f zitadel.env - - - name: test dashboard.env file gen CockroachDB - run: test -f dashboard.env - - - name: test relay.env file gen CockroachDB - run: test -f relay.env + if bash infrastructure_files/getting-started-with-zitadel.sh >stdout.txt 2>stderr.txt; then + echo "Expected the retired Zitadel installer to fail" + exit 1 + fi + test ! -s stdout.txt + grep -Fq "Zitadel support and existing Zitadel deployments are not deprecated." stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-quickstart" stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/identity-providers/zitadel" stderr.txt + grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-guide" stderr.txt + grep -Fq "removed in NetBird v0.80" stderr.txt diff --git a/infrastructure_files/getting-started-with-dex.sh b/infrastructure_files/getting-started-with-dex.sh index 5e605f19c..a9be19eb3 100755 --- a/infrastructure_files/getting-started-with-dex.sh +++ b/infrastructure_files/getting-started-with-dex.sh @@ -1,557 +1,19 @@ -#!/bin/bash +#!/usr/bin/env bash -set -e +cat >&2 <<'EOF' +ERROR: This legacy installation script has been retired and no longer runs. -# NetBird Getting Started with Dex IDP -# This script sets up NetBird with Dex as the identity provider +Dex support is not deprecated. For new deployments, use getting-started.sh: -# Sed pattern to strip base64 padding characters -SED_STRIP_PADDING='s/=//g' +https://docs.netbird.io/selfhosted/selfhosted-quickstart -check_docker_compose() { - if command -v docker-compose &> /dev/null - then - echo "docker-compose" - return - fi - if docker compose --help &> /dev/null - then - echo "docker compose" - return - fi +The current installer includes NetBird's embedded Dex-based identity provider. +Local users and external identity providers can be managed through the +NetBird Dashboard: - echo "docker-compose is not installed or not in PATH. Please follow the steps from the official guide: https://docs.docker.com/engine/install/" > /dev/stderr - exit 1 -} +https://docs.netbird.io/selfhosted/identity-providers/local -check_jq() { - if ! command -v jq &> /dev/null - then - echo "jq is not installed or not in PATH, please install with your package manager. e.g. sudo apt install jq" > /dev/stderr - exit 1 - fi - return 0 -} - -get_main_ip_address() { - if [[ "$OSTYPE" == "darwin"* ]]; then - interface=$(route -n get default | grep 'interface:' | awk '{print $2}') - ip_address=$(ifconfig "$interface" | grep 'inet ' | awk '{print $2}') - else - interface=$(ip route | grep default | awk '{print $5}' | head -n 1) - ip_address=$(ip addr show "$interface" | grep 'inet ' | awk '{print $2}' | cut -d'/' -f1) - fi - - echo "$ip_address" - return 0 -} - -check_nb_domain() { - DOMAIN=$1 - if [[ "$DOMAIN-x" == "-x" ]]; then - echo "The NETBIRD_DOMAIN variable cannot be empty." > /dev/stderr - return 1 - fi - - if [[ "$DOMAIN" == "netbird.example.com" ]]; then - echo "The NETBIRD_DOMAIN cannot be netbird.example.com" > /dev/stderr - return 1 - fi - return 0 -} - -read_nb_domain() { - READ_NETBIRD_DOMAIN="" - echo -n "Enter the domain you want to use for NetBird (e.g. netbird.my-domain.com): " > /dev/stderr - read -r READ_NETBIRD_DOMAIN < /dev/tty - if ! check_nb_domain "$READ_NETBIRD_DOMAIN"; then - read_nb_domain - fi - echo "$READ_NETBIRD_DOMAIN" - return 0 -} - -get_turn_external_ip() { - TURN_EXTERNAL_IP_CONFIG="#external-ip=" - IP=$(curl -s -4 https://jsonip.com | jq -r '.ip') - if [[ "x-$IP" != "x-" ]]; then - TURN_EXTERNAL_IP_CONFIG="external-ip=$IP" - fi - echo "$TURN_EXTERNAL_IP_CONFIG" - return 0 -} - -wait_dex() { - set +e - echo -n "Waiting for Dex to become ready (via $NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN)" - counter=1 - while true; do - # Check Dex through Caddy proxy (also validates TLS is working) - if curl -sk -f -o /dev/null "$NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN/dex/.well-known/openid-configuration" 2>/dev/null; then - break - fi - if [[ $counter -eq 60 ]]; then - echo "" - echo "Taking too long. Checking logs..." - $DOCKER_COMPOSE_COMMAND logs --tail=20 caddy - $DOCKER_COMPOSE_COMMAND logs --tail=20 dex - fi - echo -n " ." - sleep 2 - counter=$((counter + 1)) - done - echo " done" - set -e - return 0 -} - -init_environment() { - CADDY_SECURE_DOMAIN="" - NETBIRD_PORT=80 - NETBIRD_HTTP_PROTOCOL="http" - NETBIRD_RELAY_PROTO="rel" - TURN_USER="self" - TURN_PASSWORD=$(openssl rand -base64 32 | sed "$SED_STRIP_PADDING") - NETBIRD_RELAY_AUTH_SECRET=$(openssl rand -base64 32 | sed "$SED_STRIP_PADDING") - TURN_MIN_PORT=49152 - TURN_MAX_PORT=65535 - TURN_EXTERNAL_IP_CONFIG=$(get_turn_external_ip) - - # Generate secrets for Dex - DEX_DASHBOARD_CLIENT_SECRET=$(openssl rand -base64 32 | sed "$SED_STRIP_PADDING") - - # Generate admin password - NETBIRD_ADMIN_PASSWORD=$(openssl rand -base64 16 | sed "$SED_STRIP_PADDING") - - if ! check_nb_domain "$NETBIRD_DOMAIN"; then - NETBIRD_DOMAIN=$(read_nb_domain) - fi - - if [[ "$NETBIRD_DOMAIN" == "use-ip" ]]; then - NETBIRD_DOMAIN=$(get_main_ip_address) - else - NETBIRD_PORT=443 - CADDY_SECURE_DOMAIN=", $NETBIRD_DOMAIN:$NETBIRD_PORT" - NETBIRD_HTTP_PROTOCOL="https" - NETBIRD_RELAY_PROTO="rels" - fi - - check_jq - - DOCKER_COMPOSE_COMMAND=$(check_docker_compose) - - if [[ -f dex.yaml ]]; then - echo "Generated files already exist, if you want to reinitialize the environment, please remove them first." - echo "You can use the following commands:" - echo " $DOCKER_COMPOSE_COMMAND down --volumes # to remove all containers and volumes" - echo " rm -f docker-compose.yml Caddyfile dex.yaml dashboard.env turnserver.conf management.json relay.env" - echo "Be aware that this will remove all data from the database, and you will have to reconfigure the dashboard." - exit 1 - fi - - echo Rendering initial files... - render_docker_compose > docker-compose.yml - render_caddyfile > Caddyfile - render_dex_config > dex.yaml - render_dashboard_env > dashboard.env - render_management_json > management.json - render_turn_server_conf > turnserver.conf - render_relay_env > relay.env - - echo -e "\nStarting Dex IDP\n" - $DOCKER_COMPOSE_COMMAND up -d caddy dex - - # Wait for Dex to be ready (through caddy proxy) - sleep 3 - wait_dex - - echo -e "\nStarting NetBird services\n" - $DOCKER_COMPOSE_COMMAND up -d - - echo -e "\nDone!\n" - echo "You can access the NetBird dashboard at $NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN" - echo "" - echo "Login with the following credentials:" - install -m 600 /dev/null .env - printf 'Email: admin@%s\nPassword: %s\n' \ - "$NETBIRD_DOMAIN" "$NETBIRD_ADMIN_PASSWORD" >> .env - echo "Email: admin@$NETBIRD_DOMAIN" - echo "Password: $NETBIRD_ADMIN_PASSWORD" - echo "" - echo "Dex admin UI is not available (Dex has no built-in UI)." - echo "To add more users, edit dex.yaml and restart: $DOCKER_COMPOSE_COMMAND restart dex" - return 0 -} - -render_caddyfile() { - cat < /dev/null; then - ADMIN_PASSWORD_HASH=$(htpasswd -bnBC 10 "" "$NETBIRD_ADMIN_PASSWORD" | tr -d ':\n') - elif command -v python3 &> /dev/null; then - ADMIN_PASSWORD_HASH=$(python3 -c "import bcrypt; print(bcrypt.hashpw('$NETBIRD_ADMIN_PASSWORD'.encode(), bcrypt.gensalt(rounds=10)).decode())" 2>/dev/null || echo "") - fi - - # Fallback to a known hash if we can't generate one - if [[ -z "$ADMIN_PASSWORD_HASH" ]]; then - # This is hash of "password" - user should change it - ADMIN_PASSWORD_HASH='$2a$10$2b2cU8CPhOTaGrs1HRQuAueS7JTT5ZHsHSzYiFPm1leZck7Mc8T4W' - NETBIRD_ADMIN_PASSWORD="password" - echo "Warning: Could not generate password hash. Using default password: password. Please change it in dex.yaml" > /dev/stderr - fi - - cat </dev/null || cat /proc/sys/kernel/random/uuid 2>/dev/null || echo "admin-user-id-001")" - -# Optional: Add external identity provider connectors -# connectors: -# - type: github -# id: github -# name: GitHub -# config: -# clientID: \$GITHUB_CLIENT_ID -# clientSecret: \$GITHUB_CLIENT_SECRET -# redirectURI: $NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN/dex/callback -# -# - type: ldap -# id: ldap -# name: LDAP -# config: -# host: ldap.example.com:636 -# insecureNoSSL: false -# bindDN: cn=admin,dc=example,dc=com -# bindPW: admin -# userSearch: -# baseDN: ou=users,dc=example,dc=com -# filter: "(objectClass=person)" -# username: uid -# idAttr: uid -# emailAttr: mail -# nameAttr: cn -EOF - return 0 -} - -render_turn_server_conf() { - cat <&2 <<'EOF' +ERROR: This legacy installation script has been retired and no longer runs. -handle_request_command_status() { - PARSED_RESPONSE=$1 - FUNCTION_NAME=$2 - RESPONSE=$3 - if [[ $PARSED_RESPONSE -ne 0 ]]; then - echo "ERROR calling $FUNCTION_NAME:" $(echo "$RESPONSE" | jq -r '.message') > /dev/stderr - exit 1 - fi -} +Zitadel support and existing Zitadel deployments are not deprecated. -handle_zitadel_request_response() { - PARSED_RESPONSE=$1 - FUNCTION_NAME=$2 - RESPONSE=$3 - if [[ $PARSED_RESPONSE == "null" ]]; then - echo "ERROR calling $FUNCTION_NAME:" $(echo "$RESPONSE" | jq -r '.message') > /dev/stderr - exit 1 - fi - sleep 1 -} +For new deployments, use getting-started.sh: -check_docker_compose() { - if command -v docker-compose &> /dev/null - then - echo "docker-compose" - return - fi - if docker compose --help &> /dev/null - then - echo "docker compose" - return - fi +https://docs.netbird.io/selfhosted/selfhosted-quickstart - echo "docker-compose is not installed or not in PATH. Please follow the steps from the official guide: https://docs.docker.com/engine/install/" > /dev/stderr - exit 1 -} +The current installer includes NetBird's embedded Dex-based identity provider. +Zitadel can be added as an external identity provider directly through the +NetBird Dashboard: -check_jq() { - if ! command -v jq &> /dev/null - then - echo "jq is not installed or not in PATH, please install with your package manager. e.g. sudo apt install jq" > /dev/stderr - exit 1 - fi -} +https://docs.netbird.io/selfhosted/identity-providers/zitadel -wait_crdb() { - set +e - while true; do - if $DOCKER_COMPOSE_COMMAND exec -T zdb curl -sf -o /dev/null 'http://localhost:8080/health?ready=1'; then - break - fi - echo -n " ." - sleep 5 - done - echo " done" - set -e -} +Standalone Zitadel and other custom identity-provider deployments remain +supported through the advanced guide: -init_crdb() { - if [[ $ZITADEL_DATABASE == "cockroach" ]]; then - echo -e "\nInitializing Zitadel's CockroachDB\n\n" - $DOCKER_COMPOSE_COMMAND up -d zdb - echo "" - # shellcheck disable=SC2028 - echo -n "Waiting CockroachDB to become ready" - wait_crdb - $DOCKER_COMPOSE_COMMAND exec -T zdb /bin/bash -c "cp /cockroach/certs/* /zitadel-certs/ && cockroach cert create-client --overwrite --certs-dir /zitadel-certs/ --ca-key /zitadel-certs/ca.key zitadel_user && chown -R 1000:1000 /zitadel-certs/" - handle_request_command_status $? "init_crdb failed" "" - fi -} +https://docs.netbird.io/selfhosted/selfhosted-guide -get_main_ip_address() { - if [[ "$OSTYPE" == "darwin"* ]]; then - interface=$(route -n get default | grep 'interface:' | awk '{print $2}') - ip_address=$(ifconfig "$interface" | grep 'inet ' | awk '{print $2}') - else - interface=$(ip route | grep default | awk '{print $5}' | head -n 1) - ip_address=$(ip addr show "$interface" | grep 'inet ' | awk '{print $2}' | cut -d'/' -f1) - fi - - echo "$ip_address" -} - -wait_pat() { - PAT_PATH=$1 - set +e - while true; do - if [[ -f "$PAT_PATH" ]]; then - break - fi - echo -n " ." - sleep 1 - done - echo " done" - set -e -} - -wait_api() { - INSTANCE_URL=$1 - PAT=$2 - set +e - counter=1 - while true; do - FLAGS="-s" - if [[ $counter -eq 45 ]]; then - FLAGS="-v" - echo "" - fi - - curl $FLAGS --fail --connect-timeout 1 -o /dev/null "$INSTANCE_URL/auth/v1/users/me" -H "Authorization: Bearer $PAT" - if [[ $? -eq 0 ]]; then - break - fi - if [[ $counter -eq 45 ]]; then - echo "" - echo "Unable to connect to Zitadel for more than 45s, please check the output above, your firewall rules and the caddy container logs to confirm if there are any issues provisioning TLS certificates" - fi - echo -n " ." - sleep 1 - counter=$((counter + 1)) - done - echo " done" - set -e -} - -create_new_project() { - INSTANCE_URL=$1 - PAT=$2 - PROJECT_NAME="NETBIRD" - - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/management/v1/projects" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{"name": "'"$PROJECT_NAME"'"}' - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.id') - handle_zitadel_request_response "$PARSED_RESPONSE" "create_new_project" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -create_new_application() { - INSTANCE_URL=$1 - PAT=$2 - APPLICATION_NAME=$3 - BASE_REDIRECT_URL1=$4 - BASE_REDIRECT_URL2=$5 - LOGOUT_URL=$6 - ZITADEL_DEV_MODE=$7 - DEVICE_CODE=$8 - - if [[ $DEVICE_CODE == "true" ]]; then - GRANT_TYPES='["OIDC_GRANT_TYPE_AUTHORIZATION_CODE","OIDC_GRANT_TYPE_DEVICE_CODE","OIDC_GRANT_TYPE_REFRESH_TOKEN"]' - else - GRANT_TYPES='["OIDC_GRANT_TYPE_AUTHORIZATION_CODE","OIDC_GRANT_TYPE_REFRESH_TOKEN"]' - fi - - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/management/v1/projects/$PROJECT_ID/apps/oidc" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{ - "name": "'"$APPLICATION_NAME"'", - "redirectUris": [ - "'"$BASE_REDIRECT_URL1"'", - "'"$BASE_REDIRECT_URL2"'" - ], - "postLogoutRedirectUris": [ - "'"$LOGOUT_URL"'" - ], - "RESPONSETypes": [ - "OIDC_RESPONSE_TYPE_CODE" - ], - "grantTypes": '"$GRANT_TYPES"', - "appType": "OIDC_APP_TYPE_USER_AGENT", - "authMethodType": "OIDC_AUTH_METHOD_TYPE_NONE", - "version": "OIDC_VERSION_1_0", - "devMode": '"$ZITADEL_DEV_MODE"', - "accessTokenType": "OIDC_TOKEN_TYPE_JWT", - "accessTokenRoleAssertion": true, - "skipNativeAppSuccessPage": true - }' - ) - - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.clientId') - handle_zitadel_request_response "$PARSED_RESPONSE" "create_new_application" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -create_service_user() { - INSTANCE_URL=$1 - PAT=$2 - - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/management/v1/users/machine" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{ - "userName": "netbird-service-account", - "name": "Netbird Service Account", - "description": "Netbird Service Account for IDP management", - "accessTokenType": "ACCESS_TOKEN_TYPE_JWT" - }' - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.userId') - handle_zitadel_request_response "$PARSED_RESPONSE" "create_service_user" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -create_service_user_secret() { - INSTANCE_URL=$1 - PAT=$2 - USER_ID=$3 - - RESPONSE=$( - curl -sS -X PUT "$INSTANCE_URL/management/v1/users/$USER_ID/secret" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{}' - ) - SERVICE_USER_CLIENT_ID=$(echo "$RESPONSE" | jq -r '.clientId') - handle_zitadel_request_response "$SERVICE_USER_CLIENT_ID" "create_service_user_secret_id" "$RESPONSE" - SERVICE_USER_CLIENT_SECRET=$(echo "$RESPONSE" | jq -r '.clientSecret') - handle_zitadel_request_response "$SERVICE_USER_CLIENT_SECRET" "create_service_user_secret" "$RESPONSE" -} - -add_organization_user_manager() { - INSTANCE_URL=$1 - PAT=$2 - USER_ID=$3 - - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/management/v1/orgs/me/members" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{ - "userId": "'"$USER_ID"'", - "roles": [ - "ORG_USER_MANAGER" - ] - }' - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.details.creationDate') - handle_zitadel_request_response "$PARSED_RESPONSE" "add_organization_user_manager" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -create_admin_user() { - INSTANCE_URL=$1 - PAT=$2 - USERNAME=$3 - PASSWORD=$4 - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/management/v1/users/human/_import" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{ - "userName": "'"$USERNAME"'", - "profile": { - "firstName": "Zitadel", - "lastName": "Admin" - }, - "email": { - "email": "'"$USERNAME"'", - "isEmailVerified": true - }, - "password": "'"$PASSWORD"'", - "passwordChangeRequired": true - }' - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.userId') - handle_zitadel_request_response "$PARSED_RESPONSE" "create_admin_user" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -add_instance_admin() { - INSTANCE_URL=$1 - PAT=$2 - USER_ID=$3 - - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/admin/v1/members" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{ - "userId": "'"$USER_ID"'", - "roles": [ - "IAM_OWNER" - ] - }' - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.details.creationDate') - handle_zitadel_request_response "$PARSED_RESPONSE" "add_instance_admin" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -delete_auto_service_user() { - INSTANCE_URL=$1 - PAT=$2 - - RESPONSE=$( - curl -sS -X GET "$INSTANCE_URL/auth/v1/users/me" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - ) - USER_ID=$(echo "$RESPONSE" | jq -r '.user.id') - handle_zitadel_request_response "$USER_ID" "delete_auto_service_user_get_user" "$RESPONSE" - - RESPONSE=$( - curl -sS -X DELETE "$INSTANCE_URL/admin/v1/members/$USER_ID" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.details.changeDate') - handle_zitadel_request_response "$PARSED_RESPONSE" "delete_auto_service_user_remove_instance_permissions" "$RESPONSE" - - RESPONSE=$( - curl -sS -X DELETE "$INSTANCE_URL/management/v1/orgs/me/members/$USER_ID" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.details.changeDate') - handle_zitadel_request_response "$PARSED_RESPONSE" "delete_auto_service_user_remove_org_permissions" "$RESPONSE" - echo "$PARSED_RESPONSE" -} - -delete_default_zitadel_admin() { - INSTANCE_URL=$1 - PAT=$2 - - # Search for the default zitadel-admin user - RESPONSE=$( - curl -sS -X POST "$INSTANCE_URL/management/v1/users/_search" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - -d '{ - "queries": [ - { - "userNameQuery": { - "userName": "zitadel-admin@", - "method": "TEXT_QUERY_METHOD_STARTS_WITH" - } - } - ] - }' - ) - - DEFAULT_ADMIN_ID=$(echo "$RESPONSE" | jq -r '.result[0].id // empty') - - if [ -n "$DEFAULT_ADMIN_ID" ] && [ "$DEFAULT_ADMIN_ID" != "null" ]; then - echo "Found default zitadel-admin user with ID: $DEFAULT_ADMIN_ID" - - RESPONSE=$( - curl -sS -X DELETE "$INSTANCE_URL/management/v1/users/$DEFAULT_ADMIN_ID" \ - -H "Authorization: Bearer $PAT" \ - -H "Content-Type: application/json" \ - ) - PARSED_RESPONSE=$(echo "$RESPONSE" | jq -r '.details.changeDate // "deleted"') - handle_zitadel_request_response "$PARSED_RESPONSE" "delete_default_zitadel_admin" "$RESPONSE" - - else - echo "Default zitadel-admin user not found: $RESPONSE" - fi -} - -init_zitadel() { - echo -e "\nInitializing Zitadel with NetBird's applications\n" - INSTANCE_URL="$NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN" - - TOKEN_PATH=./machinekey/zitadel-admin-sa.token - - echo -n "Waiting for Zitadel's PAT to be created " - wait_pat "$TOKEN_PATH" - echo "Reading Zitadel PAT" - PAT=$(cat $TOKEN_PATH) - if [ "$PAT" = "null" ]; then - echo "Failed requesting getting Zitadel PAT" - exit 1 - fi - - echo -n "Waiting for Zitadel to become ready " - wait_api "$INSTANCE_URL" "$PAT" - - echo "Deleting default zitadel-admin user..." - delete_default_zitadel_admin "$INSTANCE_URL" "$PAT" - - # create the zitadel project - echo "Creating new zitadel project" - PROJECT_ID=$(create_new_project "$INSTANCE_URL" "$PAT") - - ZITADEL_DEV_MODE=false - BASE_REDIRECT_URL=$NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN - if [[ $NETBIRD_HTTP_PROTOCOL == "http" ]]; then - ZITADEL_DEV_MODE=true - fi - - # create zitadel spa applications - echo "Creating new Zitadel SPA Dashboard application" - DASHBOARD_APPLICATION_CLIENT_ID=$(create_new_application "$INSTANCE_URL" "$PAT" "Dashboard" "$BASE_REDIRECT_URL/nb-auth" "$BASE_REDIRECT_URL/nb-silent-auth" "$BASE_REDIRECT_URL/" "$ZITADEL_DEV_MODE" "false") - - echo "Creating new Zitadel SPA Cli application" - CLI_APPLICATION_CLIENT_ID=$(create_new_application "$INSTANCE_URL" "$PAT" "Cli" "http://localhost:53000/" "http://localhost:54000/" "http://localhost:53000/" "true" "true") - - MACHINE_USER_ID=$(create_service_user "$INSTANCE_URL" "$PAT") - - SERVICE_USER_CLIENT_ID="null" - SERVICE_USER_CLIENT_SECRET="null" - - create_service_user_secret "$INSTANCE_URL" "$PAT" "$MACHINE_USER_ID" - - DATE=$(add_organization_user_manager "$INSTANCE_URL" "$PAT" "$MACHINE_USER_ID") - - ZITADEL_ADMIN_USERNAME="admin@$NETBIRD_DOMAIN" - ZITADEL_ADMIN_PASSWORD="$(openssl rand -base64 32 | sed 's/=//g')@" - - HUMAN_USER_ID=$(create_admin_user "$INSTANCE_URL" "$PAT" "$ZITADEL_ADMIN_USERNAME" "$ZITADEL_ADMIN_PASSWORD") - - DATE="null" - - DATE=$(add_instance_admin "$INSTANCE_URL" "$PAT" "$HUMAN_USER_ID") - - DATE="null" - DATE=$(delete_auto_service_user "$INSTANCE_URL" "$PAT") - if [ "$DATE" = "null" ]; then - echo "Failed deleting auto service user" - echo "Please remove it manually" - fi - - export NETBIRD_AUTH_CLIENT_ID=$DASHBOARD_APPLICATION_CLIENT_ID - export NETBIRD_AUTH_CLIENT_ID_CLI=$CLI_APPLICATION_CLIENT_ID - export NETBIRD_IDP_MGMT_CLIENT_ID=$SERVICE_USER_CLIENT_ID - export NETBIRD_IDP_MGMT_CLIENT_SECRET=$SERVICE_USER_CLIENT_SECRET - export ZITADEL_ADMIN_USERNAME - export ZITADEL_ADMIN_PASSWORD -} - -check_nb_domain() { - DOMAIN=$1 - if [ "$DOMAIN-x" == "-x" ]; then - echo "The NETBIRD_DOMAIN variable cannot be empty." > /dev/stderr - return 1 - fi - - if [ "$DOMAIN" == "netbird.example.com" ]; then - echo "The NETBIRD_DOMAIN cannot be netbird.example.com" > /dev/stderr - return 1 - fi - return 0 -} - -read_nb_domain() { - READ_NETBIRD_DOMAIN="" - echo -n "Enter the domain you want to use for NetBird (e.g. netbird.my-domain.com): " > /dev/stderr - read -r READ_NETBIRD_DOMAIN < /dev/tty - if ! check_nb_domain "$READ_NETBIRD_DOMAIN"; then - read_nb_domain - fi - echo "$READ_NETBIRD_DOMAIN" -} - -get_turn_external_ip() { - TURN_EXTERNAL_IP_CONFIG="#external-ip=" - IP=$(curl -s -4 https://jsonip.com | jq -r '.ip') - if [[ "x-$IP" != "x-" ]]; then - TURN_EXTERNAL_IP_CONFIG="external-ip=$IP" - fi - echo "$TURN_EXTERNAL_IP_CONFIG" -} - -initEnvironment() { - CADDY_SECURE_DOMAIN="" - ZITADEL_EXTERNALSECURE="false" - ZITADEL_TLS_MODE="disabled" - ZITADEL_MASTERKEY="$(openssl rand -base64 32 | head -c 32)" - NETBIRD_PORT=80 - NETBIRD_HTTP_PROTOCOL="http" - NETBIRD_RELAY_PROTO="rel" - TURN_USER="self" - TURN_PASSWORD=$(openssl rand -base64 32 | sed 's/=//g') - NETBIRD_RELAY_AUTH_SECRET=$(openssl rand -base64 32 | sed 's/=//g') - TURN_MIN_PORT=49152 - TURN_MAX_PORT=65535 - TURN_EXTERNAL_IP_CONFIG=$(get_turn_external_ip) - - if ! check_nb_domain "$NETBIRD_DOMAIN"; then - NETBIRD_DOMAIN=$(read_nb_domain) - fi - - if [ "$NETBIRD_DOMAIN" == "use-ip" ]; then - NETBIRD_DOMAIN=$(get_main_ip_address) - else - ZITADEL_EXTERNALSECURE="true" - ZITADEL_TLS_MODE="external" - NETBIRD_PORT=443 - CADDY_SECURE_DOMAIN=", $NETBIRD_DOMAIN:$NETBIRD_PORT" - NETBIRD_HTTP_PROTOCOL="https" - NETBIRD_RELAY_PROTO="rels" - fi - - if [[ "$OSTYPE" == "darwin"* ]]; then - ZIDATE_TOKEN_EXPIRATION_DATE=$(date -u -v+30M "+%Y-%m-%dT%H:%M:%SZ") - else - ZIDATE_TOKEN_EXPIRATION_DATE=$(date -u -d "+30 minutes" "+%Y-%m-%dT%H:%M:%SZ") - fi - - check_jq - - DOCKER_COMPOSE_COMMAND=$(check_docker_compose) - - if [ -f zitadel.env ]; then - echo "Generated files already exist, if you want to reinitialize the environment, please remove them first." - echo "You can use the following commands:" - echo " $DOCKER_COMPOSE_COMMAND down --volumes # to remove all containers and volumes" - echo " rm -f docker-compose.yml Caddyfile zitadel.env dashboard.env machinekey/zitadel-admin-sa.token turnserver.conf management.json relay.env" - echo "Be aware that this will remove all data from the database, and you will have to reconfigure the dashboard." - exit 1 - fi - - if [[ $ZITADEL_DATABASE == "cockroach" ]]; then - echo "Use CockroachDB as Zitadel database." - ZDB=$(renderDockerComposeCockroachDB) - ZITADEL_DB_ENV=$(renderZitadelCockroachDBEnv) - else - echo "Use Postgres as default Zitadel database." - echo "For using CockroachDB please the environment variable 'export ZITADEL_DATABASE=cockroach'." - POSTGRES_ROOT_PASSWORD="$(openssl rand -base64 32 | sed 's/=//g')@" - POSTGRES_ZITADEL_PASSWORD="$(openssl rand -base64 32 | sed 's/=//g')@" - ZDB=$(renderDockerComposePostgres) - ZITADEL_DB_ENV=$(renderZitadelPostgresEnv) - renderPostgresEnv > zdb.env - fi - - echo Rendering initial files... - renderDockerCompose > docker-compose.yml - renderCaddyfile > Caddyfile - renderZitadelEnv > zitadel.env - echo "" > dashboard.env - echo "" > turnserver.conf - echo "" > management.json - echo "" > relay.env - - mkdir -p machinekey - chmod 777 machinekey - - init_crdb - - echo -e "\nStarting Zitadel IDP for user management\n\n" - $DOCKER_COMPOSE_COMMAND up -d caddy zitadel - init_zitadel - - echo -e "\nRendering NetBird files...\n" - renderTurnServerConf > turnserver.conf - renderManagementJson > management.json - renderDashboardEnv > dashboard.env - renderRelayEnv > relay.env - - echo -e "\nStarting NetBird services\n" - $DOCKER_COMPOSE_COMMAND up -d - echo -e "\nDone!\n" - echo "You can access the NetBird dashboard at $NETBIRD_HTTP_PROTOCOL://$NETBIRD_DOMAIN" - echo "Login with the following credentials:" - install -m 600 /dev/null .env - printf 'Username: %s\nPassword: %s\n' \ - "$ZITADEL_ADMIN_USERNAME" "$ZITADEL_ADMIN_PASSWORD" >> .env - echo "Username: $ZITADEL_ADMIN_USERNAME" - echo "Password: $ZITADEL_ADMIN_PASSWORD" -} - -renderCaddyfile() { - cat < Date: Wed, 29 Jul 2026 20:52:18 +0900 Subject: [PATCH 20/32] [client] Fix UI crash on Windows builds without dark-mode support (#6958) --- client/ui/services/windowtheme_windows.go | 14 ++++++++++++++ 1 file changed, 14 insertions(+) create mode 100644 client/ui/services/windowtheme_windows.go diff --git a/client/ui/services/windowtheme_windows.go b/client/ui/services/windowtheme_windows.go new file mode 100644 index 000000000..7dbc1164b --- /dev/null +++ b/client/ui/services/windowtheme_windows.go @@ -0,0 +1,14 @@ +package services + +import "github.com/wailsapp/wails/v3/pkg/w32" + +// Wails assigns w32.AllowDarkModeForWindow only on builds >= 18334 but calls it +// without a nil check when a window requests the Dark theme, crashing older +// builds such as Windows Server 2019 (17763). Those builds still get a dark +// title bar via the pre-20H1 DWM attribute that w32.SetTheme applies, so a +// no-op stub keeps the Dark theme fully working there. +func init() { + if w32.AllowDarkModeForWindow == nil { + w32.AllowDarkModeForWindow = func(w32.HWND, bool) uintptr { return 0 } + } +} From 0b7e6a9f46ea9ea5a93bbb80463b387059e7fcd0 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Wed, 29 Jul 2026 16:54:01 +0200 Subject: [PATCH 21/32] [signal] make pprof configurable (#6963) --- signal/cmd/run.go | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/signal/cmd/run.go b/signal/cmd/run.go index 81e9cc926..a36623c6b 100644 --- a/signal/cmd/run.go +++ b/signal/cmd/run.go @@ -10,6 +10,7 @@ import ( "net/http" // nolint:gosec _ "net/http/pprof" + "os" "time" "go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc" @@ -195,12 +196,14 @@ var ( ) func startPprof() { - go func() { - log.Debugf("Starting pprof server on 127.0.0.1:6060") - if err := http.ListenAndServe("127.0.0.1:6060", nil); err != nil { - log.Fatalf("pprof server failed: %v", err) - } - }() + if pprofAddr := os.Getenv("NB_PPROF_ADDR"); pprofAddr != "" { + log.Infof("pprof enabled, listening on: %s", pprofAddr) + go func() { + if err := http.ListenAndServe(pprofAddr, nil); err != nil { + log.Fatalf("pprof server failed: %v", err) + } + }() + } } func getTLSConfigurations() ([]grpc.ServerOption, *autocert.Manager, *tls.Config, error) { From 0f5d2d91fb3d969e30cb13396e5b6917a8f6617b Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Thu, 30 Jul 2026 03:32:26 +0900 Subject: [PATCH 22/32] [client] Authorize daemon IPC callers by their local identity (#6967) --- CONTRIBUTING.md | 14 +- client/cmd/daemon_error.go | 66 ++++ client/cmd/logout.go | 2 +- client/cmd/root.go | 16 +- client/cmd/service.go | 13 +- client/cmd/service_controller.go | 177 ++++++--- client/cmd/service_json_gateway.go | 112 +++++- client/cmd/service_json_gateway_test.go | 261 +++++++++++++ client/cmd/service_params.go | 8 + client/cmd/service_pipe_other.go | 14 + client/cmd/service_pipe_windows.go | 41 +++ client/cmd/service_socket.go | 12 +- client/cmd/up.go | 6 +- client/internal/daemonaddr/owner.go | 15 + client/internal/daemonaddr/owner_unix.go | 40 ++ client/internal/daemonaddr/owner_unix_test.go | 62 ++++ client/internal/daemonaddr/owner_windows.go | 42 +++ client/internal/daemonaddr/pipe.go | 103 ++++++ client/internal/daemonaddr/pipe_other.go | 15 + client/internal/daemonaddr/pipe_test.go | 30 ++ client/internal/daemonaddr/pipe_windows.go | 59 +++ .../internal/daemonaddr/resolve_pipe_other.go | 9 + .../daemonaddr/resolve_pipe_windows.go | 82 +++++ client/internal/ipcauth/creds_stub.go | 31 ++ client/internal/ipcauth/creds_unix.go | 56 +++ client/internal/ipcauth/creds_windows.go | 194 ++++++++++ client/internal/ipcauth/forward.go | 272 ++++++++++++++ client/internal/ipcauth/forward_test.go | 214 +++++++++++ client/internal/ipcauth/identity.go | 127 +++++++ client/internal/ipcauth/peercred_bsd.go | 43 +++ client/internal/ipcauth/peercred_linux.go | 39 ++ client/internal/ipcauth/pipeserver_windows.go | 87 +++++ client/internal/ipcauth/privileged.go | 125 +++++++ client/internal/ipcauth/privileged_test.go | 134 +++++++ client/internal/ipcauth/self_unix.go | 17 + client/internal/ipcauth/self_windows.go | 35 ++ client/internal/profilemanager/config.go | 7 + client/server/login_gate_test.go | 127 +++++++ client/server/server.go | 215 +++++++++-- client/server/setconfig_mdm_test.go | 6 +- client/server/setconfig_test.go | 7 +- client/server/ssh_gate.go | 282 ++++++++++++++ client/server/ssh_gate_test.go | 348 ++++++++++++++++++ client/ssh/client/client.go | 13 +- client/ssh/proxy/proxy.go | 7 +- .../src/components/CopyToClipboard.tsx | 7 +- .../frontend/src/contexts/SettingsContext.tsx | 26 +- client/ui/frontend/src/hooks/usePrivilege.ts | 32 ++ client/ui/frontend/src/lib/errors.ts | 27 +- .../src/modules/error/ErrorDialog.tsx | 37 +- .../src/modules/settings/SettingsSSH.tsx | 86 ++++- client/ui/grpc.go | 18 +- client/ui/i18n/locales/en/common.json | 12 + client/ui/main.go | 2 +- client/ui/services/errors.go | 39 ++ client/ui/services/settings.go | 74 +++- client/ui/services/windowmanager.go | 17 +- go.mod | 4 +- 58 files changed, 3799 insertions(+), 167 deletions(-) create mode 100644 client/cmd/daemon_error.go create mode 100644 client/cmd/service_json_gateway_test.go create mode 100644 client/cmd/service_pipe_other.go create mode 100644 client/cmd/service_pipe_windows.go create mode 100644 client/internal/daemonaddr/owner.go create mode 100644 client/internal/daemonaddr/owner_unix.go create mode 100644 client/internal/daemonaddr/owner_unix_test.go create mode 100644 client/internal/daemonaddr/owner_windows.go create mode 100644 client/internal/daemonaddr/pipe.go create mode 100644 client/internal/daemonaddr/pipe_other.go create mode 100644 client/internal/daemonaddr/pipe_test.go create mode 100644 client/internal/daemonaddr/pipe_windows.go create mode 100644 client/internal/daemonaddr/resolve_pipe_other.go create mode 100644 client/internal/daemonaddr/resolve_pipe_windows.go create mode 100644 client/internal/ipcauth/creds_stub.go create mode 100644 client/internal/ipcauth/creds_unix.go create mode 100644 client/internal/ipcauth/creds_windows.go create mode 100644 client/internal/ipcauth/forward.go create mode 100644 client/internal/ipcauth/forward_test.go create mode 100644 client/internal/ipcauth/identity.go create mode 100644 client/internal/ipcauth/peercred_bsd.go create mode 100644 client/internal/ipcauth/peercred_linux.go create mode 100644 client/internal/ipcauth/pipeserver_windows.go create mode 100644 client/internal/ipcauth/privileged.go create mode 100644 client/internal/ipcauth/privileged_test.go create mode 100644 client/internal/ipcauth/self_unix.go create mode 100644 client/internal/ipcauth/self_windows.go create mode 100644 client/server/login_gate_test.go create mode 100644 client/server/ssh_gate.go create mode 100644 client/server/ssh_gate_test.go create mode 100644 client/ui/frontend/src/hooks/usePrivilege.ts diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 261083783..d9c0b416e 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -234,12 +234,22 @@ cd client/ui task dev ``` -Pass daemon flags after `--`: +Pass daemon flags after `--`, pointing the UI at the socket the daemon serves: ``` -task dev -- --daemon-addr=tcp://127.0.0.1:41731 +task dev -- --daemon-addr=unix:///var/run/netbird.sock # Linux, macOS +task dev -- --daemon-addr=npipe://netbird # Windows ``` +On Windows the daemon serves a named pipe (`npipe://netbird`). Which path that +ends up being depends on what the daemon may create: as a service or elevated it +serves `\\.\pipe\ProtectedPrefix\Administrators\netbird`, which no unprivileged +process can take from it, and otherwise it falls back to `\\.\pipe\netbird`. +Clients try both and check who owns the pipe before using the plain one. Avoid +`tcp://127.0.0.1:41731`: loopback TCP carries no caller identity, so the daemon +refuses the operations that require an administrator and you will not exercise +those paths. + Production build (frontend assets embedded into the binary, output in `client/ui/bin/`): ``` diff --git a/client/cmd/daemon_error.go b/client/cmd/daemon_error.go new file mode 100644 index 000000000..0d5b1307e --- /dev/null +++ b/client/cmd/daemon_error.go @@ -0,0 +1,66 @@ +package cmd + +import ( + "errors" + "fmt" + "strings" + + "google.golang.org/genproto/googleapis/rpc/errdetails" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// daemonCallError prepares a daemon error for display. A refusal the daemon +// raised because the operation needs root/administrator is already guidance +// written for the user, so it is surfaced on its own instead of buried under the +// gRPC envelope and the name of the RPC that hit it. Anything else is wrapped +// with context as usual. +func daemonCallError(context string, err error) error { + if guidance, ok := privilegeGuidance(err); ok { + return errors.New(guidance) + } + return fmt.Errorf("%s: %w", context, err) +} + +// privilegeGuidance renders the daemon's privilege refusal as a summary and the +// command that performs the operation with the privileges it needs. It reports +// false for any other error. +func privilegeGuidance(err error) (string, bool) { + info, ok := privilegeErrorInfo(err) + if !ok { + return "", false + } + + summary := info.GetMetadata()[ipcauth.ErrorMetaSummary] + command := info.GetMetadata()[ipcauth.ErrorMetaCommand] + if summary == "" { + // Detail without a summary: fall back to the status message, which + // carries the same text. + summary = strings.TrimSpace(gstatus.Convert(err).Message()) + } + if command == "" { + return summary, true + } + + return fmt.Sprintf("%s\n\n %s\n", summary, command), true +} + +// privilegeErrorInfo returns the daemon's privilege-refusal detail, if the error +// carries one. +func privilegeErrorInfo(err error) (*errdetails.ErrorInfo, bool) { + if err == nil { + return nil, false + } + + for _, detail := range gstatus.Convert(err).Details() { + info, ok := detail.(*errdetails.ErrorInfo) + if !ok { + continue + } + if info.GetReason() == ipcauth.ErrorReasonPrivilegeRequired && info.GetDomain() == ipcauth.ErrorDomain { + return info, true + } + } + return nil, false +} diff --git a/client/cmd/logout.go b/client/cmd/logout.go index 1a5281acb..dcd7b5075 100644 --- a/client/cmd/logout.go +++ b/client/cmd/logout.go @@ -46,7 +46,7 @@ var logoutCmd = &cobra.Command{ } if _, err := daemonClient.Logout(ctx, req); err != nil { - return fmt.Errorf("deregister: %v", err) + return daemonCallError("deregister", err) } cmd.Println("Deregistered successfully") diff --git a/client/cmd/root.go b/client/cmd/root.go index f1ef32717..ebaae7e3e 100644 --- a/client/cmd/root.go +++ b/client/cmd/root.go @@ -20,7 +20,6 @@ import ( "github.com/spf13/cobra" "github.com/spf13/pflag" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" daddr "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/profilemanager" @@ -91,6 +90,7 @@ var ( // Don't resolve for service commands — they create the socket, not connect to it. if !isServiceCmd(cmd) { daemonAddr = daddr.ResolveUnixDaemonAddr(daemonAddr) + daemonAddr = daddr.ResolveDaemonAddr(daemonAddr) } return nil }, @@ -143,10 +143,10 @@ func init() { defaultDaemonAddr := "unix:///var/run/netbird.sock" if runtime.GOOS == "windows" { - defaultDaemonAddr = "tcp://127.0.0.1:41731" + defaultDaemonAddr = daddr.WindowsPipeAddr } - rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp]://[path|host:port]") + rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp|npipe]://[path|host:port|name]") rootCmd.PersistentFlags().StringVarP(&managementURL, "management-url", "m", "", fmt.Sprintf("Management Service URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultManagementURL)) rootCmd.PersistentFlags().StringVar(&adminURL, "admin-url", "", fmt.Sprintf("Admin Panel URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultAdminURL)) rootCmd.PersistentFlags().StringVarP(&logLevel, "log-level", "l", "info", "sets NetBird log level") @@ -269,12 +269,10 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e ctx, cancel := context.WithTimeout(ctx, time.Second*10) defer cancel() - return grpc.DialContext( - ctx, - strings.TrimPrefix(addr, "tcp://"), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpc.WithBlock(), - ) + target, opts := daddr.DialTarget(addr) + opts = append(opts, grpc.WithBlock()) + + return grpc.DialContext(ctx, target, opts...) } // WithBackOff execute function in backoff cycle. diff --git a/client/cmd/service.go b/client/cmd/service.go index b0a56c71a..7410d60ea 100644 --- a/client/cmd/service.go +++ b/client/cmd/service.go @@ -33,10 +33,15 @@ var ( ) type program struct { - ctx context.Context - cancel context.CancelFunc - serv *grpc.Server - jsonServ *http.Server + ctx context.Context + cancel context.CancelFunc + serv *grpc.Server + jsonServ *http.Server + // jsonClient is the gateway's own connection to the daemon. It is held so + // shutting the gateway down also closes it: nothing else references it once + // the handlers are registered, so its transport goroutines would otherwise + // outlive the server. + jsonClient *grpc.ClientConn jsonServMu sync.Mutex serverInstance *server.Server serverInstanceMu sync.Mutex diff --git a/client/cmd/service_controller.go b/client/cmd/service_controller.go index 5ef13a0a6..9ba3bce25 100644 --- a/client/cmd/service_controller.go +++ b/client/cmd/service_controller.go @@ -5,6 +5,7 @@ package cmd import ( "context" "fmt" + "runtime" "time" "github.com/kardianos/service" @@ -13,6 +14,8 @@ import ( "github.com/spf13/cobra" "google.golang.org/grpc" + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/server" "github.com/netbirdio/netbird/client/system" @@ -26,6 +29,31 @@ func validateJSONSocketFlags() error { return nil } +// daemonServerOptions installs the transport credentials that expose each +// caller's kernel-authenticated identity to the handlers, which is what lets +// the daemon require root/administrator for privileged operations. +// +// The handshake exchanges no bytes, so older CLI and UI binaries still +// interoperate. Callers on a TCP socket carry no identity at all: the daemon +// keeps serving them, and the privileged operations deny them, so a warning is +// logged to make the loss of functionality visible. +func daemonServerOptions(network string) []grpc.ServerOption { + if network == "tcp" { + log.Warnf("daemon is listening on TCP (%s): callers carry no verifiable identity over TCP, "+ + "so privileged operations (SSH root login, SSH auth, enabling the SSH server, management URL changes, "+ + "deregistration) will be denied. Use a unix socket, or npipe:// on Windows", daemonAddr) + return nil + } + + creds := ipcauth.NewTransportCredentials() + if creds == nil { + log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS) + return nil + } + + return []grpc.ServerOption{grpc.Creds(creds)} +} + func (p *program) Start(svc service.Service) error { // Start should not block. Do the actual work async. log.Info("starting NetBird service") //nolint @@ -37,68 +65,106 @@ func (p *program) Start(svc service.Service) error { // Collect static system and platform information system.UpdateStaticInfoAsync() - // in any case, even if configuration does not exists we run daemon to serve CLI gRPC API. - p.serv = grpc.NewServer() - - daemonListener, err := listenOnAddress(daemonAddr) - if err != nil { - return fmt.Errorf("listen daemon interface: %w", err) + // A daemon installed before named-pipe support has the loopback TCP address + // persisted. Move it to the named pipe so an upgraded daemon can identify + // its callers instead of silently serving an unauthenticated socket. + if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok { + log.Infof("daemon address %q predates named-pipe support, listening on %q so callers can be identified", daemonAddr, migrated) + daemonAddr = migrated } - var jsonListener *socketListener - if enableJSONSocket { - jsonListener, err = listenOnAddress(jsonSocket) - if err != nil { - _ = daemonListener.Close() - return fmt.Errorf("listen daemon JSON interface: %w", err) - } - } else { - removeStaleUnixSocketForAddress(jsonSocket) + network, _, err := parseListenAddress(daemonAddr) + if err != nil { + return fmt.Errorf("parse daemon address: %w", err) + } + + // in any case, even if configuration does not exists we run daemon to serve CLI gRPC API. + p.serv = grpc.NewServer(daemonServerOptions(network)...) + + daemonListener, jsonListener, err := listenDaemonSockets() + if err != nil { + return err } go func() { - defer daemonListener.Close() - if jsonListener != nil { - defer jsonListener.Close() - } - - if err := daemonListener.chmodUnixSocket("daemon"); err != nil { - log.Error(err) - return - } - if jsonListener != nil { - if err := jsonListener.chmodUnixSocket("daemon JSON"); err != nil { - log.Error(err) - return - } - } - - serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled) - if err := serverInstance.Start(); err != nil { - log.Fatalf("failed to start daemon: %v", err) - } - proto.RegisterDaemonServiceServer(p.serv, serverInstance) - - p.serverInstanceMu.Lock() - p.serverInstance = serverInstance - p.serverInstanceMu.Unlock() - - if jsonListener != nil { - if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil { - log.Fatalf("failed to start daemon JSON server: %v", err) - } - } else { - log.Debug("daemon JSON socket disabled") - } - - log.Printf("started daemon server: %v", daemonListener.address) - if err := p.serv.Serve(daemonListener.Listener); err != nil { - log.Errorf("failed to serve daemon requests: %v", err) + // Fatal here rather than inside serve, so serve's deferred listener + // closes run before the process exits. + if err := p.serve(daemonListener, jsonListener); err != nil { + log.Fatalf("failed to %v", err) } }() return nil } +// listenDaemonSockets opens the daemon control socket and, when it is enabled, the +// JSON gateway socket. The control socket is closed again if the second one fails, +// so a failed start leaves nothing listening. The returned JSON listener is nil +// when the socket is disabled. +func listenDaemonSockets() (*socketListener, *socketListener, error) { + daemonListener, err := listenOnAddress(daemonAddr) + if err != nil { + return nil, nil, fmt.Errorf("listen daemon interface: %w", err) + } + + if !enableJSONSocket { + removeStaleUnixSocketForAddress(jsonSocket) + return daemonListener, nil, nil + } + + jsonListener, err := listenOnAddress(jsonSocket) + if err != nil { + if cerr := daemonListener.Close(); cerr != nil { + log.Debugf("close daemon listener: %v", cerr) + } + return nil, nil, fmt.Errorf("listen daemon JSON interface: %w", err) + } + + return daemonListener, jsonListener, nil +} + +// serve brings up the daemon server on an already-open control socket and blocks +// until it stops. jsonListener is nil when the JSON socket is disabled. A returned +// error means the daemon cannot run at all and the caller is expected to exit; the +// failures it recovers from on its own are logged here. +func (p *program) serve(daemonListener, jsonListener *socketListener) error { + defer daemonListener.Close() + if jsonListener != nil { + defer jsonListener.Close() + } + + // chmodUnixSocket is a no-op for a nil listener and for a non-unix one. + if err := daemonListener.chmodUnixSocket("daemon"); err != nil { + log.Error(err) + return nil + } + if err := jsonListener.chmodUnixSocket("daemon JSON"); err != nil { + log.Error(err) + return nil + } + + serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled) + if err := serverInstance.Start(); err != nil { + return fmt.Errorf("start daemon: %w", err) + } + proto.RegisterDaemonServiceServer(p.serv, serverInstance) + + p.serverInstanceMu.Lock() + p.serverInstance = serverInstance + p.serverInstanceMu.Unlock() + + if jsonListener == nil { + log.Debug("daemon JSON socket disabled") + } else if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil { + return fmt.Errorf("start daemon JSON server: %w", err) + } + + log.Printf("started daemon server: %v", daemonListener.address) + if err := p.serv.Serve(daemonListener.Listener); err != nil { + log.Errorf("failed to serve daemon requests: %v", err) + } + return nil +} + func (p *program) Stop(srv service.Service) error { p.serverInstanceMu.Lock() if p.serverInstance != nil { @@ -113,8 +179,13 @@ func (p *program) Stop(srv service.Service) error { p.cancel() p.jsonServMu.Lock() - jsonServ := p.jsonServ + jsonServ, jsonClient := p.jsonServ, p.jsonClient p.jsonServMu.Unlock() + if jsonClient != nil { + if err := jsonClient.Close(); err != nil { + log.Debugf("close daemon JSON gateway client: %v", err) + } + } if jsonServ != nil { shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 2*time.Second) if err := jsonServ.Shutdown(shutdownCtx); err != nil { diff --git a/client/cmd/service_json_gateway.go b/client/cmd/service_json_gateway.go index 29c1a6456..b6864f338 100644 --- a/client/cmd/service_json_gateway.go +++ b/client/cmd/service_json_gateway.go @@ -5,27 +5,123 @@ package cmd import ( "context" "errors" + "fmt" "net" "net/http" - "strings" + "sync" "time" "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" log "github.com/sirupsen/logrus" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" ) -func grpcGatewayEndpoint(addr string) string { - return strings.TrimPrefix(addr, "tcp://") +// jsonPeerIdentity is the context key under which the connecting HTTP client's +// identity is stashed for the lifetime of its connection. +type jsonPeerIdentity struct{} + +// jsonPeerIdentityValue pairs the identity with whether it could be read at +// all, so an unreadable identity is forwarded as "unknown" rather than omitted. +type jsonPeerIdentityValue struct { + id ipcauth.Identity + known bool +} + +// jsonConnContext reads the identity of the client connecting to the JSON +// socket and stashes it on the connection's context. The gateway re-dials the +// daemon in-process, so the daemon would otherwise see every JSON request as +// coming from the daemon itself. +func jsonConnContext(ctx context.Context, c net.Conn) context.Context { + value := jsonPeerIdentityValue{} + id, err := ipcauth.ConnIdentity(c) + if err != nil { + log.Warnf("json gateway: cannot read HTTP client identity, privileged operations will be denied for this connection: %v", err) + } else { + value.id = id + value.known = true + } + return context.WithValue(ctx, jsonPeerIdentity{}, value) +} + +// forwardIdentity stamps the HTTP client's identity onto every call the gateway +// makes to the daemon. +// +// It is an interceptor on the gateway's client connection rather than a +// runtime.WithMetadata annotator because grpc-gateway skips annotators when no +// request header maps to metadata, which an HTTP/1.0 request with no Host header +// over a unix socket achieves. The daemon would then receive no marker, see its own +// identity as the transport peer, and authorize the request as the daemon itself. +// An interceptor runs for every RPC whatever the request looked like. +func forwardIdentity(ctx context.Context) context.Context { + value, ok := ctx.Value(jsonPeerIdentity{}).(jsonPeerIdentityValue) + if !ok { + // No ConnContext ran for this request, so forward an unknown identity: + // the daemon must not mistake its own identity for the client's. + return ipcauth.WithForwardedIdentity(ctx, ipcauth.Identity{}, false) + } + return ipcauth.WithForwardedIdentity(ctx, value.id, value.known) +} + +func forwardIdentityUnary(ctx context.Context, method string, req, reply any, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error { + return invoker(forwardIdentity(ctx), method, req, reply, cc, opts...) +} + +func forwardIdentityStream(ctx context.Context, desc *grpc.StreamDesc, cc *grpc.ClientConn, method string, streamer grpc.Streamer, opts ...grpc.CallOption) (grpc.ClientStream, error) { + return streamer(forwardIdentity(ctx), desc, cc, method, opts...) +} + +// reservedHeaderWarning limits the dropped-header warning to the first occurrence. +var reservedHeaderWarning sync.Once + +// jsonIncomingHeaderMatcher keeps an HTTP client from supplying the metadata the +// gateway uses to forward its identity. grpc-gateway turns "Grpc-Metadata-" +// headers into gRPC metadata and joins them ahead of what its annotators add, so +// without this filter a JSON client could send its own x-netbird-fwd-uid and the +// daemon would authorize that instead of the client's real identity. +func jsonIncomingHeaderMatcher(key string) (string, bool) { + mapped, ok := runtime.DefaultHeaderMatcher(key) + if !ok { + return "", false + } + if ipcauth.IsReservedForwardKey(mapped) { + // Warn once: any client can send these on every request, so warning each + // time hands it a way to fill the log. The rest are debug-level. + reservedHeaderWarning.Do(func() { + log.Warnf("json gateway: dropping reserved header %q from a request: only the gateway may set the caller's identity", key) + }) + log.Debugf("json gateway: dropping reserved header %q", key) + return "", false + } + return mapped, true } func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint string) error { - mux := runtime.NewServeMux() - opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())} - if err := proto.RegisterDaemonServiceHandlerFromEndpoint(p.ctx, mux, grpcGatewayEndpoint(daemonEndpoint), opts); err != nil { + if jsonListener.network == "tcp" { + log.Warnf("daemon JSON socket is listening on TCP (%s): callers carry no verifiable identity over TCP, "+ + "so privileged operations will be denied for JSON clients", jsonListener.address) + } + + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + + // grpc.NewClient does not connect until the first request, so registering + // the handler here cannot block daemon startup. + target, opts := daemonaddr.DialTarget(daemonEndpoint) + opts = append(opts, + grpc.WithChainUnaryInterceptor(forwardIdentityUnary), + grpc.WithChainStreamInterceptor(forwardIdentityStream), + ) + conn, err := grpc.NewClient(target, opts...) + if err != nil { + return fmt.Errorf("create daemon client for JSON gateway: %w", err) + } + if err := proto.RegisterDaemonServiceHandler(p.ctx, mux, conn); err != nil { + if cerr := conn.Close(); cerr != nil { + log.Debugf("close daemon client after failed JSON gateway registration: %v", cerr) + } return err } @@ -35,10 +131,12 @@ func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint BaseContext: func(net.Listener) context.Context { return p.ctx }, + ConnContext: jsonConnContext, } p.jsonServMu.Lock() p.jsonServ = jsonServer + p.jsonClient = conn p.jsonServMu.Unlock() go func() { diff --git a/client/cmd/service_json_gateway_test.go b/client/cmd/service_json_gateway_test.go new file mode 100644 index 000000000..dfeef1c46 --- /dev/null +++ b/client/cmd/service_json_gateway_test.go @@ -0,0 +1,261 @@ +//go:build !windows && !ios && !android + +package cmd + +import ( + "context" + "net" + "net/http" + "path/filepath" + "testing" + "time" + + "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/peer" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// The JSON gateway runs inside the daemon and re-dials it locally, so every JSON +// request reaches a handler with the daemon's own identity as the transport peer. +// The gateway therefore forwards its HTTP client's identity as metadata, and the +// daemon authorizes that instead of itself. These tests drive the real wiring +// (jsonConnContext, forwardIdentity, jsonIncomingHeaderMatcher) and check the +// identity a handler would end up authorizing. + +// daemonSideCtx is what a handler sees for a gateway-relayed call. The transport +// peer must be this process's own identity: the gateway is the daemon, so the two +// cannot differ, and hardcoding root here instead would describe a state that +// never occurs. +func daemonSideCtx(t *testing.T, md metadata.MD) context.Context { + t.Helper() + self, err := ipcauth.CurrentProcessIdentity() + if err != nil { + t.Skipf("cannot read this process's identity: %v", err) + } + ctx := peer.NewContext(context.Background(), &peer.Peer{ + AuthInfo: ipcauth.AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: self, + }, + }) + return metadata.NewIncomingContext(ctx, md) +} + +// gatewayMetadata reproduces what the daemon receives for a JSON request: the +// mux annotates the context from the request's headers, then the interceptor on the +// gateway's client connection stamps the caller's identity. The order matters, +// since the interceptor must win over anything a header put there. +func gatewayMetadata(t *testing.T, req *http.Request, ctx context.Context) metadata.MD { + t.Helper() + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + annotated, err := runtime.AnnotateContext(ctx, mux, req, + "/daemon.DaemonService/SetConfig", + runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig")) + if err != nil { + t.Fatalf("annotate: %v", err) + } + + md, ok := metadata.FromOutgoingContext(forwardIdentity(annotated)) + if !ok { + t.Fatal("the interceptor produced no metadata") + } + return md +} + +// clientCtx is the connection context jsonConnContext would have produced for an +// HTTP client whose identity the gateway could read. +func clientCtx(id ipcauth.Identity, known bool) context.Context { + return context.WithValue(context.Background(), jsonPeerIdentity{}, + jsonPeerIdentityValue{id: id, known: known}) +} + +// An HTTP client must not be able to name its own identity. grpc-gateway turns +// Grpc-Metadata- headers into gRPC metadata, so without the header filter and +// the interceptor overwriting the reserved keys, this request would authorize as +// uid 0. +func TestJSONGateway_ForgedIdentityHeaderIsDropped(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Uid", "0") + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Gid", "0") + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd", "1") + req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Sid", "S-1-5-18") + + caller := ipcauth.Identity{UID: 31000, GID: 31000} + md := gatewayMetadata(t, req, clientCtx(caller, true)) + + id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)) + if !ok { + t.Fatal("the forwarded identity should be usable") + } + if id.IsPrivileged() { + t.Errorf("forged header was believed: authorized as %v", id) + } + if id.UID != caller.UID { + t.Errorf("authorized as uid %d, want the real client %d", id.UID, caller.UID) + } +} + +// A request with no headers at all (HTTP/1.0 needs no Host, and a unix socket +// yields no host:port) makes grpc-gateway produce no metadata whatsoever and skip +// its annotators: "if len(pairs) == 0 { return ctx, nil, nil }" in +// runtime/context.go. That is why the identity is stamped by an interceptor +// instead. This is the case that previously reached the gate as the daemon itself. +func TestJSONGateway_HeaderlessRequestIsStillMarkedForwarded(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + req.Header = http.Header{} + req.Host = "" + + caller := ipcauth.Identity{UID: 31000, GID: 31000} + ctx := clientCtx(caller, true) + + // Pin the skip path itself: if grpc-gateway ever produced a pair here, this + // test would still pass below while no longer covering what it was written for. + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + annotated, err := runtime.AnnotateContext(ctx, mux, req, + "/daemon.DaemonService/SetConfig", + runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig")) + if err != nil { + t.Fatalf("annotate: %v", err) + } + if md, ok := metadata.FromOutgoingContext(annotated); ok { + t.Fatalf("grpc-gateway produced metadata %v for a headerless request; "+ + "this test no longer covers the annotator-skip path", md) + } + + md := gatewayMetadata(t, req, ctx) + + id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)) + if !ok { + t.Fatal("the forwarded identity should be usable") + } + if id.UID != caller.UID || id.IsPrivileged() { + t.Errorf("authorized as %v, want the real client uid %d", id, caller.UID) + } +} + +// When the gateway cannot read its client's identity (a TCP JSON socket, say) it +// forwards the marker alone. The daemon must then report "unidentified" so the +// privileged operations refuse, rather than falling back to the gateway's own +// identity. +func TestJSONGateway_UnreadableClientIdentityIsUnidentified(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + + md := gatewayMetadata(t, req, clientCtx(ipcauth.Identity{}, false)) + + if id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)); ok { + t.Errorf("a request with no client identity was authorized as %v", id) + } +} + +// A request that never passed through jsonConnContext (no stashed identity) must +// also come out unidentified rather than as the daemon. +func TestJSONGateway_MissingConnContextIsUnidentified(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil) + if err != nil { + t.Fatal(err) + } + + md := gatewayMetadata(t, req, context.Background()) + + if id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)); ok { + t.Errorf("a request with no connection context was authorized as %v", id) + } +} + +// End to end over a real unix socket: the gateway reads the connecting client's +// identity from the socket itself, so a client cannot present anything else. +func TestJSONGateway_IdentityComesFromTheSocket(t *testing.T) { + mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher)) + + type observed struct { + md metadata.MD + } + seen := make(chan observed, 1) + + srv := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ctx, err := runtime.AnnotateContext(r.Context(), mux, r, + "/daemon.DaemonService/SetConfig", + runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig")) + if err != nil { + t.Errorf("annotate: %v", err) + return + } + md, _ := metadata.FromOutgoingContext(forwardIdentity(ctx)) + seen <- observed{md: md} + w.WriteHeader(http.StatusOK) + }), + ReadHeaderTimeout: 5 * time.Second, + ConnContext: jsonConnContext, + } + + sock := filepath.Join(t.TempDir(), "http.sock") + ln, err := net.Listen("unix", sock) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := srv.Close(); err != nil { + t.Logf("close server: %v", err) + } + }) + go func() { + if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed { + t.Logf("serve: %v", err) + } + }() + + conn, err := net.Dial("unix", sock) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := conn.Close(); err != nil { + t.Logf("close conn: %v", err) + } + }) + + // Forge the identity headers on the wire as well. + request := "POST /daemon.DaemonService/SetConfig HTTP/1.1\r\n" + + "Host: localhost\r\n" + + "Grpc-Metadata-X-Netbird-Fwd: 1\r\n" + + "Grpc-Metadata-X-Netbird-Fwd-Uid: 0\r\n" + + "Content-Length: 0\r\n\r\n" + if _, err := conn.Write([]byte(request)); err != nil { + t.Fatal(err) + } + + select { + case got := <-seen: + self, err := ipcauth.CurrentProcessIdentity() + if err != nil { + t.Skipf("cannot read this process's identity: %v", err) + } + // The socket peer is this test process, so that is the identity the + // gateway must forward, not the uid 0 the request asked for. + if uids := got.md.Get("x-netbird-fwd-uid"); len(uids) != 1 { + t.Fatalf("x-netbird-fwd-uid = %v, want exactly the gateway's own value", uids) + } + id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, got.md)) + if !ok { + t.Fatal("the forwarded identity should be usable") + } + if id.UID != self.UID { + t.Errorf("authorized as uid %d, want the socket peer %d", id.UID, self.UID) + } + case <-time.After(5 * time.Second): + t.Fatal("the gateway never handled the request") + } +} diff --git a/client/cmd/service_params.go b/client/cmd/service_params.go index f25087a69..750b22ae6 100644 --- a/client/cmd/service_params.go +++ b/client/cmd/service_params.go @@ -13,6 +13,7 @@ import ( "github.com/spf13/cobra" "github.com/netbirdio/netbird/client/configs" + "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/util" ) @@ -125,6 +126,13 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) { if !rootCmd.PersistentFlags().Changed("daemon-addr") && params.DaemonAddr != "" { daemonAddr = params.DaemonAddr + // An install that predates named-pipe support has the loopback TCP + // address saved. Callers carry no identity over TCP, so move it to the + // pipe instead of restoring a socket the daemon cannot authorize on. + if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok { + cmd.Printf("Moving the saved daemon address from %s to %s so the daemon can identify its callers\n", daemonAddr, migrated) + daemonAddr = migrated + } } if !serviceCmd.PersistentFlags().Changed("json-socket") && params.JSONSocket != "" { diff --git a/client/cmd/service_pipe_other.go b/client/cmd/service_pipe_other.go new file mode 100644 index 000000000..c7cc72469 --- /dev/null +++ b/client/cmd/service_pipe_other.go @@ -0,0 +1,14 @@ +//go:build !windows + +package cmd + +import ( + "fmt" + "net" +) + +// listenNamedPipe is Windows-only: no other platform serves the daemon on a +// named pipe. +func listenNamedPipe(string) (net.Listener, string, error) { + return nil, "", fmt.Errorf("named pipes are only supported on Windows") +} diff --git a/client/cmd/service_pipe_windows.go b/client/cmd/service_pipe_windows.go new file mode 100644 index 000000000..b6e860f51 --- /dev/null +++ b/client/cmd/service_pipe_windows.go @@ -0,0 +1,41 @@ +//go:build windows + +package cmd + +import ( + "errors" + "fmt" + "net" + + "github.com/Microsoft/go-winio" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// listenNamedPipe creates the daemon control pipe and reports the path it ended +// up on. The security descriptor lets any local caller connect, as a Unix socket +// at 0666 does, and the privileged operations are authorized separately from the +// caller's token. +// +// The protected name comes first so that an unprivileged process cannot take the +// name before the service does. Creating it requires being an administrator or +// LocalSystem, so a daemon an ordinary user runs themselves, as in netstack mode, +// falls back to the plain name; clients try both and check who serves them. +func listenNamedPipe(name string) (net.Listener, string, error) { + var errs []error + for _, path := range daemonaddr.PipePaths(name) { + listener, err := winio.ListenPipe(path, &winio.PipeConfig{ + SecurityDescriptor: ipcauth.DefaultPipeSDDL(), + }) + if err != nil { + log.Debugf("not serving the daemon on %s: %v", path, err) + errs = append(errs, fmt.Errorf("%s: %w", path, err)) + continue + } + return listener, path, nil + } + + return nil, "", errors.Join(errs...) +} diff --git a/client/cmd/service_socket.go b/client/cmd/service_socket.go index f825a4062..ed1f001a7 100644 --- a/client/cmd/service_socket.go +++ b/client/cmd/service_socket.go @@ -26,6 +26,14 @@ func listenOnAddress(addr string) (*socketListener, error) { return nil, err } + if network == "npipe" { + listener, path, err := listenNamedPipe(address) + if err != nil { + return nil, err + } + return &socketListener{Listener: listener, network: network, address: path}, nil + } + if network == "unix" { removeStaleUnixSocket(address) } @@ -41,11 +49,11 @@ func listenOnAddress(addr string) (*socketListener, error) { func parseListenAddress(addr string) (string, string, error) { network, address, ok := strings.Cut(addr, "://") if !ok || network == "" || address == "" { - return "", "", fmt.Errorf("address must be in [unix|tcp]://[path|host:port] format: %q", addr) + return "", "", fmt.Errorf("address must be in [unix|tcp|npipe]://[path|host:port|name] format: %q", addr) } switch network { - case "unix", "tcp": + case "unix", "tcp", "npipe": return network, address, nil default: return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network) diff --git a/client/cmd/up.go b/client/cmd/up.go index 2d9731f26..142bcf6bd 100644 --- a/client/cmd/up.go +++ b/client/cmd/up.go @@ -325,7 +325,7 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable { log.Warnf("setConfig method is not available in the daemon: %s", st.Message()) } else { - return fmt.Errorf("call service setConfig method: %v", err) + return daemonCallError("call service setConfig method", err) } } @@ -379,7 +379,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ } if loginErr != nil { - return fmt.Errorf("login failed: %v", loginErr) + return daemonCallError("login failed", loginErr) } if loginResp.NeedsSSOLogin { @@ -392,7 +392,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ ProfileName: &profileID, Username: &username, }); err != nil { - return fmt.Errorf("call service up method: %v", err) + return daemonCallError("call service up method", err) } return nil diff --git a/client/internal/daemonaddr/owner.go b/client/internal/daemonaddr/owner.go new file mode 100644 index 000000000..c476f9ae6 --- /dev/null +++ b/client/internal/daemonaddr/owner.go @@ -0,0 +1,15 @@ +package daemonaddr + +// DaemonRunsAsSelf reports whether the daemon listening at addr runs as this very +// user. That is what makes an unprivileged daemon authorize this process for the +// changes it otherwise restricts to root or an administrator, so a client can tell +// up front whether those controls are usable instead of letting a save fail. +// +// It is answered from the ownership of the socket or pipe the daemon created, so it +// costs no round trip and needs no cooperation from the daemon. Ownership that +// cannot be read is reported as false, including for a TCP address, so a caller +// reading this as "the daemon would allow it" fails closed. The daemon remains the +// only thing that authorizes anything: this only decides what a client offers. +func DaemonRunsAsSelf(addr string) bool { + return daemonRunsAsSelf(addr) +} diff --git a/client/internal/daemonaddr/owner_unix.go b/client/internal/daemonaddr/owner_unix.go new file mode 100644 index 000000000..493e6528d --- /dev/null +++ b/client/internal/daemonaddr/owner_unix.go @@ -0,0 +1,40 @@ +//go:build !windows + +package daemonaddr + +import ( + "os" + "strings" + "syscall" + + log "github.com/sirupsen/logrus" +) + +// daemonRunsAsSelf compares the owner of the daemon's Unix socket with this +// process's uid. Root is not treated specially here: a root caller is privileged +// on its own merits, and a root-owned socket says nothing about the caller. +func daemonRunsAsSelf(addr string) bool { + path, ok := strings.CutPrefix(addr, "unix://") + if !ok { + return false + } + + info, err := os.Stat(path) + if err != nil { + log.Debugf("stat daemon socket %s: %v", path, err) + return false + } + + // Only a socket says anything about a daemon. A directory or a leftover + // regular file at that path is not one, and reading it as "the daemon runs as + // us" would offer controls the daemon then refuses. + if info.Mode()&os.ModeSocket == 0 { + return false + } + + stat, ok := info.Sys().(*syscall.Stat_t) + if !ok { + return false + } + return stat.Uid == uint32(os.Getuid()) +} diff --git a/client/internal/daemonaddr/owner_unix_test.go b/client/internal/daemonaddr/owner_unix_test.go new file mode 100644 index 000000000..363c7d95d --- /dev/null +++ b/client/internal/daemonaddr/owner_unix_test.go @@ -0,0 +1,62 @@ +//go:build !windows + +package daemonaddr + +import ( + "net" + "os" + "path/filepath" + "testing" +) + +// A socket this user created means the daemon runs as this user, which is the +// rootless case where the daemon delegates its authority to its own identity. +func TestDaemonRunsAsSelf_OwnSocket(t *testing.T) { + path := filepath.Join(t.TempDir(), "netbird.sock") + ln, err := net.Listen("unix", path) + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { + if err := ln.Close(); err != nil { + t.Logf("close listener: %v", err) + } + }) + + if !DaemonRunsAsSelf("unix://" + path) { + t.Error("a socket owned by this user must count as the daemon running as us") + } +} + +// Everything that is not a readable socket of ours has to answer false, because +// the caller reads a true as "the daemon would authorize me". +func TestDaemonRunsAsSelf_FailsClosed(t *testing.T) { + dir := t.TempDir() + + // A socket owned by another user, which is what a root-run daemon looks like + // to an unprivileged client. Only assertable when we are not root ourselves. + rootOwned := "unix:///var/run/netbird.sock" + if _, err := os.Stat("/var/run/netbird.sock"); err == nil && os.Getuid() != 0 { + if DaemonRunsAsSelf(rootOwned) { + t.Error("a socket owned by another user must not count as ours") + } + } + + for name, addr := range map[string]string{ + "missing socket": "unix://" + filepath.Join(dir, "absent.sock"), + "tcp address": "tcp://127.0.0.1:41731", + "named pipe": "npipe://netbird", + "empty": "", + "no scheme": filepath.Join(dir, "absent.sock"), + "directory": "unix://" + dir, + "unknown scheme": "http://localhost:8080", + "scheme only": "unix://", + "relative socket": "unix://netbird.sock", + } { + t.Run(name, func(t *testing.T) { + if DaemonRunsAsSelf(addr) { + t.Errorf("%q must not count as a daemon running as us", addr) + } + }) + } +} diff --git a/client/internal/daemonaddr/owner_windows.go b/client/internal/daemonaddr/owner_windows.go new file mode 100644 index 000000000..1cd2bba15 --- /dev/null +++ b/client/internal/daemonaddr/owner_windows.go @@ -0,0 +1,42 @@ +//go:build windows + +package daemonaddr + +import ( + "context" + "strings" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// daemonRunsAsSelf reads the owner of the daemon's pipe. A daemon running as the +// service account owns its pipe as LocalSystem, and an elevated one as +// BUILTIN\Administrators, so only a daemon the user started themselves matches. +func daemonRunsAsSelf(addr string) bool { + name, ok := strings.CutPrefix(addr, pipeScheme) + if !ok { + return false + } + + for _, path := range PipePaths(name) { + // Bounded: this runs on the UI's path for deciding which controls to + // offer, so a pipe that does not answer promptly must not stall it. A + // timeout leaves the caller unprivileged, which only disables controls. + ctx, cancel := context.WithTimeout(context.Background(), probeTimeout) + conn, err := dialPipe(ctx, path) + cancel() + if err != nil { + continue + } + + owned := ipcauth.PipeOwnedBySelf(conn) + if cerr := conn.Close(); cerr != nil { + log.Debugf("close daemon pipe %s after ownership check: %v", path, cerr) + } + return owned + } + + return false +} diff --git a/client/internal/daemonaddr/pipe.go b/client/internal/daemonaddr/pipe.go new file mode 100644 index 000000000..51815ef5e --- /dev/null +++ b/client/internal/daemonaddr/pipe.go @@ -0,0 +1,103 @@ +package daemonaddr + +import ( + "context" + "net" + "runtime" + "strings" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" +) + +const ( + // WindowsPipeAddr is the default daemon address on Windows. A named pipe + // carries the connecting process's token, which loopback TCP does not, so + // it is the only Windows transport on which the daemon can tell who is + // calling it. + WindowsPipeAddr = "npipe://netbird" + + // legacyWindowsAddr is the loopback-TCP address the Windows daemon used + // before named-pipe support. + legacyWindowsAddr = "tcp://127.0.0.1:41731" + + pipeScheme = "npipe://" + + // protectedPrefix is the NPFS namespace in which only LocalSystem and + // members of BUILTIN\Administrators may create a pipe. A daemon running as + // the service account creates its pipe there so that an unprivileged process + // cannot pre-create the name, which would keep the daemon from starting and + // leave callers talking to the squatter. Opening such a pipe needs no + // privilege, so unprivileged clients still reach the daemon. + protectedPrefix = `ProtectedPrefix\Administrators\` +) + +// DialTarget returns the gRPC dial target and transport options for a daemon +// address. The npipe scheme needs a context dialer because gRPC has no +// named-pipe resolver; unix and tcp are handled by gRPC itself. +func DialTarget(addr string) (string, []grpc.DialOption) { + opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())} + + if name, ok := strings.CutPrefix(addr, pipeScheme); ok { + paths := PipePaths(name) + opts = append(opts, grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { + return dialPipePaths(ctx, paths) + })) + return "passthrough:///netbird-daemon-pipe", opts + } + + return strings.TrimPrefix(addr, "tcp://"), opts +} + +// PipePath maps an npipe address name ("netbird", from "npipe://netbird") to a +// Windows named-pipe path (\\.\pipe\netbird). A fully qualified path is left as +// is. +func PipePath(name string) string { + if strings.HasPrefix(name, `\\`) { + return name + } + return `\\.\pipe\` + name +} + +// PipePaths returns the paths a daemon control pipe may live at for an npipe +// address name, in the order both sides must try them: the protected name first, +// then the plain one. +// +// The daemon serves the first it can create, which is the protected name when it +// runs as the service account and the plain one when it runs as an ordinary user, +// as it does in netstack mode. Clients therefore have to try both, and because a +// client cannot tell from the name alone who created the pipe, the plain name is +// only usable once the server's identity has been checked: see +// verifyPipeServer. +// +// A fully qualified path is what the operator asked for and is used as is. +func PipePaths(name string) []string { + if strings.HasPrefix(name, `\\`) { + return []string{name} + } + return []string{PipePath(protectedPrefix + name), PipePath(name)} +} + +// IsProtectedPipePath reports whether a pipe path is in the namespace only an +// administrator or LocalSystem can create in, which is what lets a client trust +// such a pipe from its name alone. +func IsProtectedPipePath(path string) bool { + return strings.HasPrefix(path, `\\.\pipe\`+protectedPrefix) +} + +// MigrateLegacy upgrades the pre-named-pipe Windows daemon address to the named +// pipe, reporting whether it rewrote the address. Existing installs persist the +// daemon address, so without this an upgraded daemon would keep listening on +// loopback TCP, where callers carry no identity and privileged operations would +// have to be refused for everyone. Only the exact legacy default is rewritten: +// a deliberately chosen custom address is left alone. +func MigrateLegacy(addr string) (string, bool) { + return migrateLegacyForOS(runtime.GOOS, addr) +} + +func migrateLegacyForOS(goos, addr string) (string, bool) { + if goos == "windows" && addr == legacyWindowsAddr { + return WindowsPipeAddr, true + } + return addr, false +} diff --git a/client/internal/daemonaddr/pipe_other.go b/client/internal/daemonaddr/pipe_other.go new file mode 100644 index 000000000..04e8e7331 --- /dev/null +++ b/client/internal/daemonaddr/pipe_other.go @@ -0,0 +1,15 @@ +//go:build !windows + +package daemonaddr + +import ( + "context" + "fmt" + "net" +) + +// dialPipePaths is Windows-only: no other platform serves the daemon on a named +// pipe. +func dialPipePaths(context.Context, []string) (net.Conn, error) { + return nil, fmt.Errorf("named pipes are only supported on Windows") +} diff --git a/client/internal/daemonaddr/pipe_test.go b/client/internal/daemonaddr/pipe_test.go new file mode 100644 index 000000000..b9dfd90f1 --- /dev/null +++ b/client/internal/daemonaddr/pipe_test.go @@ -0,0 +1,30 @@ +package daemonaddr + +import ( + "slices" + "testing" +) + +// The protected name must be tried before the plain one on both sides: it is the +// one an unprivileged process cannot create, so preferring it is what keeps a +// squatter from owning the name the service daemon would otherwise use. +func TestPipePaths_PrefersTheProtectedName(t *testing.T) { + got := PipePaths("netbird") + want := []string{ + `\\.\pipe\ProtectedPrefix\Administrators\netbird`, + `\\.\pipe\netbird`, + } + if !slices.Equal(got, want) { + t.Errorf("PipePaths = %q, want %q", got, want) + } +} + +// An operator who passes a full path chose exactly one pipe, so neither side may +// look anywhere else. +func TestPipePaths_QualifiedPathIsUsedAsIs(t *testing.T) { + path := `\\.\pipe\custom-netbird` + got := PipePaths(path) + if !slices.Equal(got, []string{path}) { + t.Errorf("PipePaths = %q, want just %q", got, path) + } +} diff --git a/client/internal/daemonaddr/pipe_windows.go b/client/internal/daemonaddr/pipe_windows.go new file mode 100644 index 000000000..3cd10a6c3 --- /dev/null +++ b/client/internal/daemonaddr/pipe_windows.go @@ -0,0 +1,59 @@ +//go:build windows + +package daemonaddr + +import ( + "context" + "errors" + "fmt" + "net" + + "github.com/Microsoft/go-winio" + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" + + "github.com/netbirdio/netbird/client/internal/ipcauth" +) + +// dialPipePaths connects to the first path that answers with a pipe server this +// client may trust, and returns the last error when none does. +func dialPipePaths(ctx context.Context, paths []string) (net.Conn, error) { + var lastErr error + for _, path := range paths { + conn, err := dialPipe(ctx, path) + if err != nil { + log.Debugf("dial daemon pipe %s: %v", path, err) + lastErr = err + continue + } + + // A pipe in the protected namespace could only have been created by an + // administrator or LocalSystem, so its name is the guarantee. Any other + // name has to be checked, because any local user can create one. + if !IsProtectedPipePath(path) { + if err := ipcauth.PipeServerTrusted(conn); err != nil { + if closeErr := conn.Close(); closeErr != nil { + log.Debugf("close untrusted pipe %s: %v", path, closeErr) + } + lastErr = fmt.Errorf("%s: %w", path, err) + continue + } + } + + return conn, nil + } + + if lastErr == nil { + lastErr = errors.New("no daemon pipe to connect to") + } + return nil, lastErr +} + +// dialPipe connects to the daemon control pipe at SECURITY_IDENTIFICATION. +// winio's plain DialPipe connects at SECURITY_ANONYMOUS, under which the daemon +// cannot read the caller's token at all. Identification lets the daemon read the +// caller's SID and groups without granting it the ability to act as the caller. +func dialPipe(ctx context.Context, path string) (net.Conn, error) { + access := uint32(windows.GENERIC_READ | windows.GENERIC_WRITE) + return winio.DialPipeAccessImpLevel(ctx, path, access, winio.PipeImpLevelIdentification) +} diff --git a/client/internal/daemonaddr/resolve_pipe_other.go b/client/internal/daemonaddr/resolve_pipe_other.go new file mode 100644 index 000000000..1aede8453 --- /dev/null +++ b/client/internal/daemonaddr/resolve_pipe_other.go @@ -0,0 +1,9 @@ +//go:build !windows + +package daemonaddr + +// ResolveDaemonAddr is a no-op off Windows, where there is no named-pipe +// default to fall back from. +func ResolveDaemonAddr(addr string) string { + return addr +} diff --git a/client/internal/daemonaddr/resolve_pipe_windows.go b/client/internal/daemonaddr/resolve_pipe_windows.go new file mode 100644 index 000000000..d12ddb15d --- /dev/null +++ b/client/internal/daemonaddr/resolve_pipe_windows.go @@ -0,0 +1,82 @@ +//go:build windows + +package daemonaddr + +import ( + "net" + "strings" + "time" + + "github.com/Microsoft/go-winio" + log "github.com/sirupsen/logrus" +) + +// probeTimeout bounds each transport probe. Both are local, so a daemon that is +// listening answers immediately and one that is not fails immediately. +const probeTimeout = 300 * time.Millisecond + +// ResolveDaemonAddr keeps a client on the named pipe and never silently moves it +// off. When the pipe does not answer it checks the legacy loopback TCP address, so +// a client meeting a daemon that has not restarted since the upgrade can say what +// is wrong, but it does not connect there. +// +// Using that address automatically would be a downgrade the user never asked for: +// any local process can bind 127.0.0.1 while the daemon is not listening, and the +// transport carries no caller identity, so a client that accepted whatever answered +// would hand a setup key, a pre-shared key or an SSO prompt to a local impostor. An +// operator who needs the legacy address during the upgrade window can still pass +// --daemon-addr explicitly, which is a deliberate choice and still refuses the +// privileged operations. +// +// Only the pipe address is resolved. A custom address is left alone, though passing +// --daemon-addr npipe://netbird explicitly is indistinguishable from the default +// here, so it is treated the same way. +func ResolveDaemonAddr(addr string) string { + if addr != WindowsPipeAddr { + return addr + } + + for _, path := range PipePaths("netbird") { + if pipeAvailable(path) { + return addr + } + } + + if tcpAvailable(legacyWindowsAddr) { + log.Warnf("the daemon is not serving %s, but something is listening on the legacy %s. "+ + "Restart the NetBird service so it serves the pipe. That address is not used automatically: "+ + "any local user can bind it and it carries no caller identity, so pass --daemon-addr %s "+ + "explicitly if you accept that", + WindowsPipeAddr, legacyWindowsAddr, legacyWindowsAddr) + } + + return addr +} + +func pipeAvailable(path string) bool { + timeout := probeTimeout + conn, err := winio.DialPipe(path, &timeout) + if err != nil { + return false + } + if err := conn.Close(); err != nil { + log.Debugf("close daemon pipe probe: %v", err) + } + return true +} + +func tcpAvailable(addr string) bool { + host := addr + if _, after, ok := strings.Cut(addr, "://"); ok { + host = after + } + + conn, err := net.DialTimeout("tcp", host, probeTimeout) + if err != nil { + return false + } + if err := conn.Close(); err != nil { + log.Debugf("close daemon TCP probe: %v", err) + } + return true +} diff --git a/client/internal/ipcauth/creds_stub.go b/client/internal/ipcauth/creds_stub.go new file mode 100644 index 000000000..154948716 --- /dev/null +++ b/client/internal/ipcauth/creds_stub.go @@ -0,0 +1,31 @@ +//go:build !linux && !darwin && !freebsd && !windows + +package ipcauth + +import ( + "errors" + "net" + + "google.golang.org/grpc/credentials" +) + +// errUnsupported is returned on platforms with no local peer-identity +// primitive, so consumers fail closed instead of guessing an identity. +var errUnsupported = errors.New("peer identity is not available on this platform") + +// NewTransportCredentials returns nil: without a peer-identity primitive the +// daemon cannot authenticate local callers, and the caller must treat that as +// "authorization cannot be enforced". +func NewTransportCredentials() credentials.TransportCredentials { + return nil +} + +// PeerIdentity always fails on this platform. +func PeerIdentity(net.Conn) (Identity, error) { + return Identity{}, errUnsupported +} + +// ConnIdentity always fails on this platform. +func ConnIdentity(net.Conn) (Identity, error) { + return Identity{}, errUnsupported +} diff --git a/client/internal/ipcauth/creds_unix.go b/client/internal/ipcauth/creds_unix.go new file mode 100644 index 000000000..688fe4623 --- /dev/null +++ b/client/internal/ipcauth/creds_unix.go @@ -0,0 +1,56 @@ +//go:build linux || darwin || freebsd + +package ipcauth + +import ( + "context" + "net" + + "google.golang.org/grpc/credentials" +) + +// NewTransportCredentials returns gRPC transport credentials that expose the +// caller's kernel-authenticated identity via IdentityFromContext. It returns +// nil on platforms that have no peer-identity primitive, which the caller must +// treat as "authorization cannot be enforced". +// +// The handshake exchanges no bytes on the wire, so a client dialing with +// insecure credentials interoperates with a server using these. That keeps +// older CLI and UI binaries working against an upgraded daemon. +func NewTransportCredentials() credentials.TransportCredentials { + return unixCreds{} +} + +// ConnIdentity extracts the caller's identity from an accepted local IPC +// connection. It is shared by the gRPC transport credentials and by the JSON +// gateway, which reads the identity of its own HTTP clients. +func ConnIdentity(conn net.Conn) (Identity, error) { + return PeerIdentity(conn) +} + +type unixCreds struct{} + +func (unixCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + return conn, AuthInfo{}, nil +} + +// ServerHandshake extracts the peer identity and fails closed when it cannot +// be read, so a connection whose caller is unknown never reaches a handler. +func (unixCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + id, err := ConnIdentity(conn) + if err != nil { + return nil, nil, err + } + return conn, AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, nil +} + +func (unixCreds) Info() credentials.ProtocolInfo { + return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()} +} + +func (unixCreds) Clone() credentials.TransportCredentials { return unixCreds{} } + +func (unixCreds) OverrideServerName(string) error { return nil } diff --git a/client/internal/ipcauth/creds_windows.go b/client/internal/ipcauth/creds_windows.go new file mode 100644 index 000000000..37f902c52 --- /dev/null +++ b/client/internal/ipcauth/creds_windows.go @@ -0,0 +1,194 @@ +//go:build windows + +package ipcauth + +import ( + "context" + "fmt" + "net" + "runtime" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" + "google.golang.org/grpc/credentials" +) + +var ( + modadvapi32 = windows.NewLazySystemDLL("advapi32.dll") + procImpersonateNamedPipeClient = modadvapi32.NewProc("ImpersonateNamedPipeClient") +) + +// DefaultPipeSDDL is the security descriptor for the daemon control pipe. +// +// D:P protected DACL, no inheritance +// (A;;GA;;;SY) allow GENERIC_ALL to LocalSystem (the daemon's service account) +// (A;;GA;;;WD) allow GENERIC_ALL to Everyone +// +// Any local caller may connect, as with a Unix socket at 0666; what a caller may +// actually do is decided from its token, not from the DACL. Remote callers are not +// a concern here: winio.ListenPipe creates the pipe with +// FILE_PIPE_REJECT_REMOTE_CLIENTS, so NPFS rejects connections from other machines +// before the descriptor is consulted. +// +// A deny ACE on the NETWORK SID would not add anything and would break callers: +// that SID is present in any network-logon token, which includes OpenSSH and WinRM +// sessions, so it denies administrators driving the CLI over SSH and denies the +// daemon itself when started from such a session. +func DefaultPipeSDDL() string { + return "D:P(A;;GA;;;SY)(A;;GA;;;WD)" +} + +// NewTransportCredentials returns gRPC transport credentials that derive the +// caller's identity from the named-pipe client token. +// +// The client must connect at SECURITY_IDENTIFICATION for the daemon to be able +// to read its token, which is what DialNamedPipe does. +func NewTransportCredentials() credentials.TransportCredentials { + return winpipeCreds{} +} + +// ConnIdentity extracts the caller's identity from an accepted named-pipe +// connection by impersonating the pipe client and reading its token. It is +// shared by the gRPC transport credentials and by the JSON gateway, which +// reads the identity of its own HTTP clients. +func ConnIdentity(conn net.Conn) (Identity, error) { + // go-winio's pipe connection embeds *win32File, which exposes Fd(). + fdConn, ok := conn.(interface{ Fd() uintptr }) + if !ok { + return Identity{}, fmt.Errorf("connection %T does not expose a pipe handle", conn) + } + return pipeClientIdentity(windows.Handle(fdConn.Fd())) +} + +type winpipeCreds struct{} + +func (winpipeCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + return conn, AuthInfo{}, nil +} + +// ServerHandshake extracts the connecting client's identity and fails closed +// when the handle or token cannot be read, so a connection whose caller is +// unknown never reaches a handler. +func (winpipeCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) { + id, err := ConnIdentity(conn) + if err != nil { + return nil, nil, err + } + return conn, AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, nil +} + +func (winpipeCreds) Info() credentials.ProtocolInfo { + return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()} +} + +func (winpipeCreds) Clone() credentials.TransportCredentials { return winpipeCreds{} } + +func (winpipeCreds) OverrideServerName(string) error { return nil } + +// pipeClientIdentity reads the connecting client's user SID, usable group +// SIDs, and elevation state by impersonating the pipe client on this thread +// and reading the resulting impersonation token. +func pipeClientIdentity(handle windows.Handle) (id Identity, err error) { + // Impersonation is per-thread, so the goroutine must stay on this thread + // until RevertToSelf, otherwise an unrelated goroutine could inherit the + // impersonated context. + runtime.LockOSThread() + + // The thread only goes back to the runtime's pool once it is provably no + // longer impersonating the client. If the revert fails, leaving it locked + // makes Go terminate it when this goroutine exits, which costs one thread + // and keeps a thread running as the client from ever being reused. + clean := false + defer func() { + if clean { + runtime.UnlockOSThread() + } + }() + + if err = impersonateNamedPipeClient(handle); err != nil { + clean = true + return Identity{}, fmt.Errorf("impersonate named pipe client: %w", err) + } + defer func() { + // Surface the revert failure only when nothing else failed: leaving + // the thread impersonated is worse than the original error. + revErr := windows.RevertToSelf() + if revErr != nil { + if err == nil { + err = fmt.Errorf("revert impersonation: %w", revErr) + } + return + } + clean = true + }() + + // openAsSelf=true opens the token with the daemon's own process context + // rather than the impersonated client's, so the open cannot fail because + // the client lacks access to its own token. + var token windows.Token + if err = windows.OpenThreadToken(windows.CurrentThread(), windows.TOKEN_QUERY, true, &token); err != nil { + return Identity{}, fmt.Errorf("open thread token: %w", err) + } + defer func() { + if cerr := token.Close(); cerr != nil { + log.Debugf("close client token: %v", cerr) + } + }() + + return identityFromToken(token) +} + +// identityFromToken reads the user SID, usable group SIDs and elevation state +// out of a Windows token. +func identityFromToken(token windows.Token) (Identity, error) { + user, err := token.GetTokenUser() + if err != nil { + return Identity{}, fmt.Errorf("read token user: %w", err) + } + + groups, err := tokenGroupSIDs(token) + if err != nil { + return Identity{}, err + } + + return Identity{ + SID: user.User.Sid.String(), + Groups: groups, + Elevated: token.IsElevated(), + }, nil +} + +// tokenGroupSIDs returns the SIDs of the groups the token can actually +// exercise. Groups that are disabled or marked deny-only are skipped: a +// UAC-filtered administrator carries BUILTIN\Administrators as deny-only, and +// treating that as membership would hand every admin account privilege it +// cannot currently use. +func tokenGroupSIDs(token windows.Token) ([]string, error) { + tg, err := token.GetTokenGroups() + if err != nil { + return nil, fmt.Errorf("read token groups: %w", err) + } + + var sids []string + for _, g := range tg.AllGroups() { + if g.Attributes&windows.SE_GROUP_ENABLED == 0 { + continue + } + if g.Attributes&windows.SE_GROUP_USE_FOR_DENY_ONLY != 0 { + continue + } + sids = append(sids, g.Sid.String()) + } + return sids, nil +} + +func impersonateNamedPipeClient(h windows.Handle) error { + r, _, e := procImpersonateNamedPipeClient.Call(uintptr(h)) + if r == 0 { + return e + } + return nil +} diff --git a/client/internal/ipcauth/forward.go b/client/internal/ipcauth/forward.go new file mode 100644 index 000000000..57749528c --- /dev/null +++ b/client/internal/ipcauth/forward.go @@ -0,0 +1,272 @@ +package ipcauth + +import ( + "context" + "crypto/rand" + "crypto/subtle" + "encoding/hex" + "fmt" + "slices" + "strconv" + "strings" + + "google.golang.org/grpc/metadata" +) + +// Metadata keys the local JSON gateway uses to forward the identity of its own +// HTTP client to the daemon. The gateway runs inside the daemon process and +// re-dials the daemon over the control socket, so without forwarding every +// JSON request would appear to come from the daemon itself. +const ( + // mdFwd marks a request as forwarded by the JSON gateway. It is always + // set, even when the gateway could not read its client's identity, so the + // daemon can tell "no identity available" apart from "not forwarded". + mdFwd = "x-netbird-fwd" + mdFwdUID = "x-netbird-fwd-uid" // Unix user ID + mdFwdGID = "x-netbird-fwd-gid" // Unix primary group ID + mdFwdSID = "x-netbird-fwd-sid" // Windows user SID + mdFwdGroup = "x-netbird-fwd-group" // Windows group SID, repeated + mdFwdElevated = "x-netbird-fwd-elevated" // Windows, "1" when elevated + + // mdFwdProof proves the forwarded identity was stamped by this process. The + // gateway runs inside the daemon, so a secret held in memory is available to + // the only legitimate producer and to nothing else. + mdFwdProof = "x-netbird-fwd-proof" +) + +// forwardKeys is every metadata key the gateway sets. An HTTP client must never +// be able to supply one itself: see IsReservedForwardKey. +var forwardKeys = []string{mdFwd, mdFwdUID, mdFwdGID, mdFwdSID, mdFwdGroup, mdFwdElevated, mdFwdProof} + +// forwardProof authenticates the gateway's forwarding metadata. It is generated +// once per daemon process and never leaves it: it is not written to disk, not +// logged, and not sent anywhere except over the daemon's own control socket to +// itself. +// +// Without it, trusting a forwarded identity rests on every layer in front of it +// stripping incoming forwarding keys, and on each key's value shape being +// distinguishable from an injected one. A single injected group SID or an +// injected "elevated" flag has the same shape as a legitimate one, so no +// cardinality rule can catch it. Requiring the proof means metadata that did not +// come from this process is refused whatever it contains. +var forwardProof = mustForwardProof() + +func mustForwardProof() string { + var buf [32]byte + if _, err := rand.Read(buf[:]); err != nil { + // Continuing would leave the forwarded path authenticated by a + // predictable value, which is worse than not starting. + panic(fmt.Sprintf("generate identity forwarding proof: %v", err)) + } + return hex.EncodeToString(buf[:]) +} + +// IsReservedForwardKey reports whether a gRPC metadata key belongs to the +// gateway's identity forwarding, and therefore must be dropped when it arrives +// from outside. +// +// grpc-gateway maps "Grpc-Metadata-" request headers into gRPC metadata and +// joins them ahead of the values its own annotators add. Without dropping these, +// an HTTP client could hand the daemon "x-netbird-fwd-uid: 0" and be believed, +// because the daemon trusts forwarded metadata when the transport peer is the +// (privileged) gateway. +func IsReservedForwardKey(key string) bool { + key = strings.ToLower(key) + return slices.Contains(forwardKeys, key) +} + +// ForwardIdentityMetadata encodes an HTTP client's identity for the JSON +// gateway to forward to the daemon. When known is false only the marker is +// set, which makes the daemon treat the caller as unidentified rather than as +// the daemon itself. +func ForwardIdentityMetadata(id Identity, known bool) metadata.MD { + md := metadata.MD{} + md.Set(mdFwd, "1") + md.Set(mdFwdProof, forwardProof) + if !known { + return md + } + + if id.IsWindows() { + md.Set(mdFwdSID, id.SID) + if len(id.Groups) > 0 { + md.Set(mdFwdGroup, id.Groups...) + } + if id.Elevated { + md.Set(mdFwdElevated, "1") + } + return md + } + + md.Set(mdFwdUID, strconv.FormatUint(uint64(id.UID), 10)) + md.Set(mdFwdGID, strconv.FormatUint(uint64(id.GID), 10)) + return md +} + +// CallerIdentity returns the identity to authorize a request against. For a +// direct connection that is the transport peer's kernel identity. For a +// request relayed by the local JSON gateway it is the identity the gateway +// forwarded, since the transport peer is then the daemon itself. +// +// A forwarded identity is only honoured when the transport peer is the daemon's +// own identity and the metadata carries this process's forwarding proof, so +// forged forwarding metadata gains a caller nothing. A forwarded request that +// carries no identity is reported as unidentified, never as the daemon. +// +// The second return value is false when no identity could be established, and +// callers MUST fail closed in that case. +func CallerIdentity(ctx context.Context) (Identity, bool) { + id, ok := IdentityFromContext(ctx) + if !ok { + return Identity{}, false + } + + // A forwarding key that arrives more than once did not come from the gateway + // alone, so nothing about the request can be trusted to describe its caller. + // Refusing outright matters because the alternative reading, "not forwarded", + // would authorize the request as the transport peer, which on the gateway's + // connection is the daemon itself. + if duplicatedForwardKey(ctx) { + return Identity{}, false + } + + forwarded := isForwarded(ctx) + + // Our own process on the other end of the socket is the JSON gateway, the only + // thing that dials the daemon from inside it. Such a call must carry a + // forwarded identity; without one there is no caller to authorize, and + // treating it as the daemon would authorize whatever reached the JSON socket. + // Only Linux reports the peer PID, so this is a belt on top of the gateway's + // interceptor rather than the sole guarantee. + if id.PID != 0 && int(id.PID) == selfPID && !forwarded { + return Identity{}, false + } + + // Only the gateway's own connection may speak for someone else. Being + // privileged is not enough and not the point: the gateway runs inside the + // daemon, so it dials as the daemon's identity whatever user that is, which + // also covers a rootless container. + if !forwarded || !IsDaemonSelf(id) { + return id, true + } + + // Speaking for someone else additionally requires the proof only this process + // holds. Refusing is the only safe reading: the transport peer here is the + // daemon itself, so falling back to it would authorize the request as the + // daemon. This is also what makes the forwarded values trustworthy once + // accepted, so they need no shape checks of their own. + if !authenticForward(ctx) { + return Identity{}, false + } + + return forwardedIdentity(ctx) +} + +// duplicatedForwardKey reports whether any forwarding key carries more than one +// value. The gateway's interceptor sets each key exactly once and replaces what +// was already there, so a repeat means a second source supplied it. +func duplicatedForwardKey(ctx context.Context) bool { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return false + } + for _, key := range forwardKeys { + // Group SIDs are legitimately repeated; the rest identify the caller. + if key == mdFwdGroup { + continue + } + if len(md.Get(key)) > 1 { + return true + } + } + return false +} + +// authenticForward reports whether the request carries this process's forwarding +// proof, which only the in-process JSON gateway can supply. +func authenticForward(ctx context.Context) bool { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return false + } + got := mdSingle(md, mdFwdProof) + return subtle.ConstantTimeCompare([]byte(got), []byte(forwardProof)) == 1 +} + +// isForwarded reports whether the request carries the JSON gateway marker. +func isForwarded(ctx context.Context) bool { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return false + } + return mdSingle(md, mdFwd) != "" +} + +// forwardedIdentity decodes the identity the JSON gateway attached. +func forwardedIdentity(ctx context.Context) (Identity, bool) { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + return Identity{}, false + } + + if sid := mdSingle(md, mdFwdSID); sid != "" { + return Identity{ + SID: sid, + // Repeated by design, one value per group, and only reachable once + // the forwarding proof has been verified. + Groups: md.Get(mdFwdGroup), + Elevated: mdSingle(md, mdFwdElevated) == "1", + }, true + } + + uid, err := strconv.ParseUint(mdSingle(md, mdFwdUID), 10, 32) + if err != nil { + return Identity{}, false + } + + id := Identity{UID: uint32(uid)} + if gid, err := strconv.ParseUint(mdSingle(md, mdFwdGID), 10, 32); err == nil { + id.GID = uint32(gid) + } + return id, true +} + +// mdSingle returns the value of a forwarded key only when exactly one was +// supplied. The gateway's interceptor sets each key exactly once, so more than one +// value means something else also supplied it, and the whole identity is treated as +// unknown rather than picking a winner. Defence in depth behind the gateway's +// header filter. +func mdSingle(md metadata.MD, key string) string { + if v := md.Get(key); len(v) == 1 { + return v[0] + } + return "" +} + +// WithForwardedIdentity stamps id onto a context's outgoing metadata for the JSON +// gateway's call to the daemon, replacing any forwarding keys already present so +// values supplied from outside cannot survive alongside it. +// +// This is deliberately not done with runtime.WithMetadata: grpc-gateway skips its +// annotators entirely when no request header maps to metadata ("if len(pairs) == 0 +// { return ctx, nil, nil }", runtime/context.go), which an HTTP/1.0 request with no +// Host header over a unix socket achieves. The daemon would then see an unmarked +// call whose transport peer is the daemon's own identity, and authorize it as the +// daemon. A client interceptor runs for every RPC regardless of headers. +func WithForwardedIdentity(ctx context.Context, id Identity, known bool) context.Context { + md, ok := metadata.FromOutgoingContext(ctx) + if !ok { + md = metadata.MD{} + } else { + md = md.Copy() + } + + for _, key := range forwardKeys { + delete(md, key) + } + for key, values := range ForwardIdentityMetadata(id, known) { + md[key] = values + } + + return metadata.NewOutgoingContext(ctx, md) +} diff --git a/client/internal/ipcauth/forward_test.go b/client/internal/ipcauth/forward_test.go new file mode 100644 index 000000000..d9adf05da --- /dev/null +++ b/client/internal/ipcauth/forward_test.go @@ -0,0 +1,214 @@ +package ipcauth + +import ( + "context" + "testing" + + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/peer" +) + +// transportCtx builds a request context as the daemon's transport credentials +// would: the identity of whoever opened the socket, plus whatever metadata the +// request carried. +func transportCtx(id Identity, md metadata.MD) context.Context { + ctx := peer.NewContext(context.Background(), &peer.Peer{ + AuthInfo: AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, + }) + if md != nil { + ctx = metadata.NewIncomingContext(ctx, md) + } + return ctx +} + +var ( + root = Identity{UID: 0} + unprivUser = Identity{UID: 1000, GID: 1000} +) + +// asDaemon pins which identity counts as this process for the duration of a test. +// Without it the test binary's own uid decides, which silently changes what +// "the gateway" means. +func asDaemon(t *testing.T, id Identity) { + t.Helper() + prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate + t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) + selfIdentity, selfKnown = id, true + selfMayDelegate = !id.IsPrivileged() +} + +func TestCallerIdentity_DirectConnections(t *testing.T) { + t.Run("no transport credentials is not an identity", func(t *testing.T) { + if _, ok := CallerIdentity(context.Background()); ok { + t.Fatal("a caller with no credentials must not be identified") + } + }) + + t.Run("a direct caller is its transport identity", func(t *testing.T) { + id, ok := CallerIdentity(transportCtx(unprivUser, nil)) + if !ok || id.UID != 1000 { + t.Fatalf("got %v ok=%t, want uid 1000", id, ok) + } + }) + + // The whole point of honouring forwarded metadata only from a privileged + // transport peer: an unprivileged caller can set any metadata it likes on its + // own connection to the daemon socket. + t.Run("an unprivileged caller cannot forge an identity", func(t *testing.T) { + asDaemon(t, root) + forged := metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdGID, "0") + id, ok := CallerIdentity(transportCtx(unprivUser, forged)) + if !ok { + t.Fatal("caller should still be identified, as itself") + } + if id.IsPrivileged() || id.UID != 1000 { + t.Fatalf("forged metadata was believed: got %v", id) + } + }) +} + +func TestCallerIdentity_GatewayForwarding(t *testing.T) { + t.Run("the gateway's client identity is used, not the gateway's own", func(t *testing.T) { + asDaemon(t, root) + md := ForwardIdentityMetadata(unprivUser, true) + id, ok := CallerIdentity(transportCtx(root, md)) + if !ok { + t.Fatal("forwarded identity should be usable") + } + if id.IsPrivileged() || id.UID != 1000 { + t.Fatalf("got %v, want the forwarded uid 1000 and not privileged", id) + } + }) + + t.Run("a privileged gateway client stays privileged", func(t *testing.T) { + asDaemon(t, root) + md := ForwardIdentityMetadata(root, true) + id, ok := CallerIdentity(transportCtx(root, md)) + if !ok || !id.IsPrivileged() { + t.Fatalf("got %v ok=%t, want a privileged identity", id, ok) + } + }) + + // A JSON socket the gateway cannot read peer credentials from (a TCP socket, + // say) must not make every request look like the daemon itself. + t.Run("an unreadable client identity is unknown, not the daemon", func(t *testing.T) { + asDaemon(t, root) + md := ForwardIdentityMetadata(Identity{}, false) + if _, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatal("a forwarded request with no identity must not be identified") + } + }) + + // grpc-gateway turns Grpc-Metadata- headers into gRPC metadata and joins + // them ahead of its annotators' values. If an HTTP client's header survived + // that, this is the shape the daemon would see: the attacker's uid 0 first, + // the real uid second. The gateway filters those headers out, and reading a + // duplicated key as unknown makes the daemon safe even if it did not. + t.Run("a duplicated key from an injected header is not believed", func(t *testing.T) { + asDaemon(t, root) + md := metadata.MD{} + md.Append(mdFwd, "1") + md.Append(mdFwdUID, "0") // injected by the HTTP client + md.Append(mdFwdUID, "1000") // appended by the gateway's annotator + if id, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatalf("injected uid was accepted: got %v", id) + } + }) + + t.Run("a duplicated marker is not believed either", func(t *testing.T) { + asDaemon(t, root) + md := metadata.MD{} + md.Append(mdFwd, "1") + md.Append(mdFwd, "1") + md.Append(mdFwdUID, "1000") + // A repeated marker must not be read as "not forwarded": that would + // authorize the request as the transport peer, which on the gateway's + // connection is the daemon itself. + if id, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatalf("a duplicated marker was believed: got %v", id) + } + }) + + // The layers in front of this (the gateway's header matcher, and its + // interceptor replacing every forwarding key) are what keep outside metadata + // from arriving at all. The proof is what the daemon can check for itself, and + // it is the only defence that works for a value whose legitimate shape is + // indistinguishable from an injected one: a lone group SID, or "elevated". + t.Run("forwarding metadata without this process's proof is refused", func(t *testing.T) { + asDaemon(t, root) + for name, md := range map[string]metadata.MD{ + "no proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0"), + "wrong proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdProof, "deadbeef"), + "windows identity without a proof": metadata.Pairs(mdFwd, "1", + mdFwdSID, "S-1-5-21-1-2-3-1001", mdFwdGroup, sidAdministrators, mdFwdElevated, "1"), + } { + t.Run(name, func(t *testing.T) { + if id, ok := CallerIdentity(transportCtx(root, md)); ok { + t.Fatalf("unstamped forwarding metadata was believed: got %v", id) + } + }) + } + }) + + // A caller that reaches the gateway cannot see the proof, so it cannot append + // a group of its own to a genuine forwarded identity: doing so would have to + // go through the interceptor, which replaces the whole set. + t.Run("a group appended to a stamped identity does not survive the interceptor", func(t *testing.T) { + asDaemon(t, root) + injected := metadata.MD{} + injected.Append(mdFwdGroup, sidAdministrators) + + ctx := WithForwardedIdentity(metadata.NewOutgoingContext(context.Background(), injected), + Identity{SID: "S-1-5-21-1-2-3-1001"}, true) + out, ok := metadata.FromOutgoingContext(ctx) + if !ok { + t.Fatal("no outgoing metadata") + } + if groups := out.Get(mdFwdGroup); len(groups) != 0 { + t.Fatalf("injected group survived: %v", groups) + } + }) +} + +func TestIsReservedForwardKey(t *testing.T) { + for _, key := range forwardKeys { + if !IsReservedForwardKey(key) { + t.Errorf("%q must be reserved", key) + } + } + + // grpc-gateway canonicalises header names, so the check has to be + // case-insensitive. + if !IsReservedForwardKey("X-Netbird-Fwd-Uid") { + t.Error("the check must be case-insensitive") + } + + for _, key := range []string{"authorization", "x-netbird", "x-netbird-fwd-uid-extra", ""} { + if IsReservedForwardKey(key) { + t.Errorf("%q must not be reserved", key) + } + } +} + +func TestForwardIdentityMetadata_AlwaysMarksForwarded(t *testing.T) { + for _, tc := range []struct { + name string + id Identity + known bool + }{ + {"known unix identity", unprivUser, true}, + {"unknown identity", Identity{}, false}, + {"windows identity", Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, true}, + } { + t.Run(tc.name, func(t *testing.T) { + md := ForwardIdentityMetadata(tc.id, tc.known) + if got := md.Get(mdFwd); len(got) != 1 || got[0] != "1" { + t.Fatalf("marker = %v, want exactly one \"1\"", got) + } + }) + } +} diff --git a/client/internal/ipcauth/identity.go b/client/internal/ipcauth/identity.go new file mode 100644 index 000000000..ff70c209a --- /dev/null +++ b/client/internal/ipcauth/identity.go @@ -0,0 +1,127 @@ +// Package ipcauth provides the kernel-authenticated identity of a local IPC +// (gRPC) caller and the transport credentials that surface it into the gRPC +// context, so the daemon can authorize individual RPCs by caller identity. +// +// On Unix the identity is read from the kernel via SO_PEERCRED (Linux) or +// LOCAL_PEERCRED (Darwin/FreeBSD). On Windows it is derived from the +// named-pipe client token. Platforms without a peer-identity primitive get no +// credentials, and every consumer must fail closed when no identity is +// available. +package ipcauth + +import ( + "context" + "fmt" + "slices" + + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" +) + +// Well-known Windows SIDs that identify a fully privileged principal. +const ( + sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM + sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE + sidNetworkService = "S-1-5-20" // NT AUTHORITY\NETWORK SERVICE + sidAdministrators = "S-1-5-32-544" // BUILTIN\Administrators +) + +// Identity is the kernel-authenticated identity of a local IPC caller. The +// zero value is not a valid identity: consumers must only use one obtained +// with a true ok/nil error return. +type Identity struct { + // UID and GID are the caller's Unix user ID and primary group ID. Both are + // zero on Windows, where SID is authoritative instead. + UID uint32 + GID uint32 + + // SID is the caller's Windows security identifier, empty on Unix. + SID string + + // Groups holds the caller's Windows group SIDs, captured from the client + // token at handshake time. Only groups that are enabled and not + // deny-only are captured, so a group listed here is one the caller can + // actually exercise. Empty on Unix. + Groups []string + + // Elevated reports whether the Windows client token is elevated (running + // as administrator, or an administrator with UAC turned off). Always false + // on Unix, where privilege is uid 0. + Elevated bool + + // PID is the caller's process ID where the platform reports it (Linux's + // SO_PEERCRED), and 0 where it does not. It identifies the daemon's own + // process dialling itself, which is what the JSON gateway does, and is never + // used to grant anything. + PID int32 +} + +// IsWindows reports whether this identity is a Windows principal (SID-based) +// rather than a Unix uid/gid principal. +func (i Identity) IsWindows() bool { + return i.SID != "" +} + +// IsPrivileged reports whether the caller is the platform's administrative +// principal, which is what the daemon requires for changes that cross the +// user-to-root boundary. +// +// On Windows the decision comes from the caller's token rather than from +// account names or group RIDs: an elevated token, one of the service accounts +// the daemon itself may run as, or a token with BUILTIN\Administrators +// enabled. A UAC-filtered administrator has that group marked deny-only, and +// deny-only groups are dropped when the identity is captured, so such a +// caller is correctly reported as unprivileged. Domain group memberships +// (Domain Admins and friends) are deliberately not consulted: they say +// nothing about what this token may do on this machine. +func (i Identity) IsPrivileged() bool { + if !i.IsWindows() { + return i.UID == 0 + } + + if i.Elevated { + return true + } + + switch i.SID { + case sidLocalSystem, sidLocalService, sidNetworkService: + return true + } + + return slices.Contains(i.Groups, sidAdministrators) +} + +// String renders the identity for audit logs and denial messages. +func (i Identity) String() string { + if i.IsWindows() { + return fmt.Sprintf("sid=%s elevated=%t", i.SID, i.Elevated) + } + return fmt.Sprintf("uid=%d gid=%d", i.UID, i.GID) +} + +// AuthInfo carries the peer Identity as a gRPC credentials.AuthInfo so +// handlers can retrieve it from the request context via IdentityFromContext. +type AuthInfo struct { + credentials.CommonAuthInfo + Identity Identity +} + +// AuthType identifies the authentication scheme. +func (AuthInfo) AuthType() string { return "netbird-ipc-peercred" } + +// IdentityFromContext extracts the caller's kernel-authenticated identity from +// the gRPC peer context. The second return value is false when no IPC +// transport credentials were negotiated, which happens on a TCP daemon socket +// and on platforms without a peer-identity primitive. Callers MUST fail closed +// in that case. +func IdentityFromContext(ctx context.Context) (Identity, bool) { + p, ok := peer.FromContext(ctx) + if !ok { + return Identity{}, false + } + info, ok := p.AuthInfo.(AuthInfo) + if !ok { + return Identity{}, false + } + return info.Identity, true +} diff --git a/client/internal/ipcauth/peercred_bsd.go b/client/internal/ipcauth/peercred_bsd.go new file mode 100644 index 000000000..6d9c5247f --- /dev/null +++ b/client/internal/ipcauth/peercred_bsd.go @@ -0,0 +1,43 @@ +//go:build darwin || freebsd + +package ipcauth + +import ( + "fmt" + "net" + + "golang.org/x/sys/unix" +) + +// PeerIdentity reads the kernel-authenticated identity of the process on the +// other end of a Unix socket via LOCAL_PEERCRED. The xucred is recorded by the +// kernel at connect() time and carries the peer's uid and its group list, of +// which the first entry is the primary group. +func PeerIdentity(conn net.Conn) (Identity, error) { + uc, ok := conn.(*net.UnixConn) + if !ok { + return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn) + } + + raw, err := uc.SyscallConn() + if err != nil { + return Identity{}, fmt.Errorf("raw conn: %w", err) + } + + var cred *unix.Xucred + var credErr error + if err := raw.Control(func(fd uintptr) { + cred, credErr = unix.GetsockoptXucred(int(fd), unix.SOL_LOCAL, unix.LOCAL_PEERCRED) + }); err != nil { + return Identity{}, fmt.Errorf("control raw conn: %w", err) + } + if credErr != nil { + return Identity{}, fmt.Errorf("read LOCAL_PEERCRED: %w", credErr) + } + + id := Identity{UID: cred.Uid} + if cred.Ngroups > 0 { + id.GID = cred.Groups[0] + } + return id, nil +} diff --git a/client/internal/ipcauth/peercred_linux.go b/client/internal/ipcauth/peercred_linux.go new file mode 100644 index 000000000..417cc1e00 --- /dev/null +++ b/client/internal/ipcauth/peercred_linux.go @@ -0,0 +1,39 @@ +//go:build linux + +package ipcauth + +import ( + "fmt" + "net" + + "golang.org/x/sys/unix" +) + +// PeerIdentity reads the kernel-authenticated identity of the process on the +// other end of a Unix socket via SO_PEERCRED. The credentials are recorded by +// the kernel at connect() time and cannot be changed for the life of the +// connection, so they are not spoofable by the caller. +func PeerIdentity(conn net.Conn) (Identity, error) { + uc, ok := conn.(*net.UnixConn) + if !ok { + return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn) + } + + raw, err := uc.SyscallConn() + if err != nil { + return Identity{}, fmt.Errorf("raw conn: %w", err) + } + + var cred *unix.Ucred + var credErr error + if err := raw.Control(func(fd uintptr) { + cred, credErr = unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED) + }); err != nil { + return Identity{}, fmt.Errorf("control raw conn: %w", err) + } + if credErr != nil { + return Identity{}, fmt.Errorf("read SO_PEERCRED: %w", credErr) + } + + return Identity{UID: cred.Uid, GID: cred.Gid, PID: cred.Pid}, nil +} diff --git a/client/internal/ipcauth/pipeserver_windows.go b/client/internal/ipcauth/pipeserver_windows.go new file mode 100644 index 000000000..7ba59d574 --- /dev/null +++ b/client/internal/ipcauth/pipeserver_windows.go @@ -0,0 +1,87 @@ +//go:build windows + +package ipcauth + +import ( + "fmt" + "net" + + log "github.com/sirupsen/logrus" + "golang.org/x/sys/windows" +) + +// PipeServerTrusted reports an error unless the pipe behind conn was created by a +// principal this client may hand secrets to. Clients call it for a pipe whose name +// carries no guarantee of its own, which is any name outside the +// ProtectedPrefix\Administrators namespace: that namespace already restricts +// creation to administrators and LocalSystem, while a plain name can be created by +// any local user before the daemon gets there. +// +// The decision is made from the pipe object's owner, not from the serving process, +// because a client cannot open a process running as another user at all, and the +// legitimate case is precisely an unprivileged client talking to a privileged +// daemon. Trusted owners are the service accounts, BUILTIN\Administrators, and +// this client's own user, the last of which is the daemon a user runs themselves +// as in netstack mode. A pipe owned by anyone else gets no setup key, pre-shared +// key or SSO prompt out of this client. +func PipeServerTrusted(conn net.Conn) error { + // go-winio's pipe connection embeds *win32File, which exposes Fd(). + fdConn, ok := conn.(interface{ Fd() uintptr }) + if !ok { + return fmt.Errorf("connection %T does not expose a pipe handle", conn) + } + + owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd())) + if err != nil { + return err + } + + if !trustedPipeOwner(owner) { + return fmt.Errorf("pipe owned by %s, which is neither an administrator nor this user", owner) + } + return nil +} + +// PipeOwnedBySelf reports whether the pipe behind conn was created by this very +// user, which is how a client recognises a daemon running as itself. Ownership it +// cannot read is reported as false. +func PipeOwnedBySelf(conn net.Conn) bool { + fdConn, ok := conn.(interface{ Fd() uintptr }) + if !ok { + return false + } + + owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd())) + if err != nil { + log.Debugf("read daemon pipe owner: %v", err) + return false + } + return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID +} + +// pipeOwnerSID reads the owner of the pipe object a client is connected to. The +// handle was opened with GENERIC_READ, which includes READ_CONTROL, so no extra +// access is needed. +func pipeOwnerSID(handle windows.Handle) (string, error) { + sd, err := windows.GetSecurityInfo(handle, windows.SE_KERNEL_OBJECT, windows.OWNER_SECURITY_INFORMATION) + if err != nil { + return "", fmt.Errorf("read pipe security info: %w", err) + } + + owner, _, err := sd.Owner() + if err != nil { + return "", fmt.Errorf("read pipe owner: %w", err) + } + return owner.String(), nil +} + +// trustedPipeOwner reports whether a pipe's owner is a principal a client may +// speak to. An elevated process's objects are owned by BUILTIN\Administrators by +// default, an unelevated one's by the user, which is why both forms appear here. +func trustedPipeOwner(owner string) bool { + switch owner { + case sidLocalSystem, sidLocalService, sidNetworkService, sidAdministrators: + return true + } + return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID +} diff --git a/client/internal/ipcauth/privileged.go b/client/internal/ipcauth/privileged.go new file mode 100644 index 000000000..95f2a50e9 --- /dev/null +++ b/client/internal/ipcauth/privileged.go @@ -0,0 +1,125 @@ +package ipcauth + +import ( + "os" + "runtime" +) + +// Fields of the ErrorInfo detail the daemon attaches to a PermissionDenied it +// raises for an operation that requires root/administrator. Clients match on +// Reason and Domain rather than on the message text, and render the summary and +// command themselves so the user gets guidance instead of a gRPC error dump. +const ( + // ErrorReasonPrivilegeRequired identifies the detail. + ErrorReasonPrivilegeRequired = "PRIVILEGE_REQUIRED" + // ErrorDomain scopes the reason to the NetBird daemon. + ErrorDomain = "daemon.netbird.io" + // ErrorMetaSummary is the one-sentence explanation of what was refused. + ErrorMetaSummary = "summary" + // ErrorMetaCommand is the command that performs the same operation with the + // privileges it needs, ready to copy and run. + ErrorMetaCommand = "command" +) + +// The identity of the process evaluating callers, captured once because it cannot +// change. selfKnown is false when it could not be read, in which case nothing is +// ever treated as this process. selfMayDelegate additionally requires this +// process to be unprivileged: see IsPrivilegedCaller. +var ( + selfIdentity Identity + selfKnown bool + selfMayDelegate bool + // selfPID is this process's PID, used to recognise the daemon dialling itself. + selfPID = os.Getpid() +) + +func init() { + id, err := CurrentProcessIdentity() + if err != nil { + return + } + selfIdentity, selfKnown = id, true + // Only an unprivileged daemon delegates its authority to its own identity. + // When it is root or LocalSystem, sharing its identity does not mean sharing + // its power: on Windows a filtered and a full token carry the same SID, so + // matching there would let a non-elevated shell of an administrator account + // act as an administrator, which is the boundary the token check exists to + // keep. + selfMayDelegate = !id.IsPrivileged() +} + +// IsDaemonSelf reports whether an identity is this very process. The JSON gateway +// runs inside the daemon and re-dials it locally, so this is what distinguishes +// the gateway from any other caller, whatever user the daemon runs as. +func IsDaemonSelf(id Identity) bool { + if !selfKnown || id.IsWindows() != selfIdentity.IsWindows() { + return false + } + if id.IsWindows() { + return id.SID != "" && id.SID == selfIdentity.SID + } + return id.UID == selfIdentity.UID +} + +// IsPrivilegedCaller reports whether an identity may make the changes the daemon +// restricts to the platform administrator. This is the daemon's own rule and +// cannot be evaluated by a client, which does not know what the daemon runs as. +// +// Beyond root/administrator it accepts a caller running as the daemon's own +// identity when the daemon is itself unprivileged. That keeps a rootless container +// working, where there is no uid 0 at all, and a Windows daemon in netstack mode, +// which needs no administrator rights. In those setups a caller sharing the +// daemon's identity can already rewrite the config files it reads and replace the +// binary it runs, so refusing it a config change would protect nothing; and an +// unprivileged daemon cannot hand out a root shell in the first place. +func IsPrivilegedCaller(id Identity) bool { + if id.IsPrivileged() { + return true + } + return selfMayDelegate && IsDaemonSelf(id) +} + +// SelfDelegatesTo returns the identity this process delegates its authority to, +// and whether it delegates at all. Only an unprivileged daemon does: see +// IsPrivilegedCaller. It exists so a refusal can name who may actually perform the +// operation, because on such a host root is neither required nor necessarily +// available. +func SelfDelegatesTo() (Identity, bool) { + if !selfKnown || !selfMayDelegate { + return Identity{}, false + } + return selfIdentity, true +} + +// PrivilegedActor names the principal a privileged operation requires, for use +// in messages shown to the user. +func PrivilegedActor() string { + if runtime.GOOS == "windows" { + return "administrator privileges" + } + return "root" +} + +// ElevatedCommand renders a command so that running it grants the privileges the +// operation needs. Windows has no in-line equivalent of sudo, so the command is +// returned unchanged and the user is expected to run it from an elevated +// terminal. +func ElevatedCommand(command string) string { + if runtime.GOOS == "windows" { + return command + } + return "sudo " + command +} + +// UpCommand renders an elevated `netbird up` with the given flags, preceded by a +// `down`. The down is what makes the command work on a connected client: `netbird +// up` prints "Already connected" and returns without applying any config flag, so +// on its own the command would appear to do nothing. It is a no-op, exit 0, when +// the client is not connected. +// +// ";" rather than "&&" so the line can be pasted into any of the shells a user +// might have: PowerShell 5.1, still the default on Windows Server, rejects "&&" +// as a syntax error. +func UpCommand(flags string) string { + return ElevatedCommand("netbird down") + "; " + ElevatedCommand("netbird up "+flags) +} diff --git a/client/internal/ipcauth/privileged_test.go b/client/internal/ipcauth/privileged_test.go new file mode 100644 index 000000000..c1c7c1543 --- /dev/null +++ b/client/internal/ipcauth/privileged_test.go @@ -0,0 +1,134 @@ +package ipcauth + +import "testing" + +// The self rule is the one place privilege is granted to something other than the +// platform administrator, so its two guards matter: it must apply only when the +// daemon is itself unprivileged, and only to a caller with the daemon's identity. +func TestIsPrivilegedCaller_SelfRule(t *testing.T) { + tests := []struct { + name string + // self stands in for the process the daemon runs as. + self Identity + selfKnown bool + caller Identity + want bool + }{ + { + name: "root is privileged whatever the daemon runs as", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{UID: 0}, + want: true, + }, + { + name: "an unprivileged daemon delegates to its own user (rootless container)", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{UID: 1000}, + want: true, + }, + { + name: "an unprivileged daemon delegates to nobody else", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{UID: 1001}, + want: false, + }, + { + // The daemon is root on a normal install, so sharing its identity is + // already covered by being root; nothing else may match. + name: "a root daemon delegates to nobody", + self: Identity{UID: 0}, + selfKnown: true, + caller: Identity{UID: 1000}, + want: false, + }, + { + // Windows netstack mode: the daemon needs no administrator rights. + name: "an unprivileged windows daemon delegates to its own SID", + self: Identity{SID: "S-1-5-21-1-2-3-1001"}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-1001"}, + want: true, + }, + { + name: "an unprivileged windows daemon delegates to no other SID", + self: Identity{SID: "S-1-5-21-1-2-3-1001"}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-1002"}, + want: false, + }, + { + // The UAC boundary: a filtered and a full token of the same account + // carry the same SID but not the same power, so an elevated daemon must + // never delegate to its own SID. + name: "an elevated windows daemon does not delegate to its own SID", + self: Identity{SID: "S-1-5-21-1-2-3-500", Elevated: true}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-500"}, + want: false, + }, + { + name: "LocalSystem is privileged on its own merits, not by delegation", + self: Identity{SID: sidLocalSystem}, + selfKnown: true, + caller: Identity{SID: sidLocalSystem}, + want: true, // LocalSystem is privileged on its own merits + }, + { + name: "identities of different kinds never match", + self: Identity{UID: 1000}, + selfKnown: true, + caller: Identity{SID: "S-1-5-21-1-2-3-1001"}, + want: false, + }, + { + name: "an unknown self identity delegates to nobody", + self: Identity{}, + selfKnown: false, + caller: Identity{UID: 1000}, + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate + t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate }) + + selfIdentity, selfKnown = tt.self, tt.selfKnown + selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged() + + if got := IsPrivilegedCaller(tt.caller); got != tt.want { + t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t", + tt.caller, tt.self, got, tt.want) + } + }) + } +} + +// The real process must never accidentally delegate: a test binary running as a +// normal user is unprivileged, so it may match itself, but nothing else. +func TestIsPrivilegedCaller_ThisProcess(t *testing.T) { + id, err := CurrentProcessIdentity() + if err != nil { + t.Skipf("cannot read this process's identity: %v", err) + } + + // This process is always allowed to act as itself: either it is privileged, or + // it is unprivileged and therefore delegates to its own identity. + if !IsPrivilegedCaller(id) { + t.Errorf("this process %v was refused its own identity", id) + } + + // A caller that is neither root nor this process must be refused, whatever + // this process happens to be. + other := Identity{UID: id.UID + 1} + if id.IsWindows() { + other = Identity{SID: id.SID + "9"} + } + if IsPrivilegedCaller(other) { + t.Errorf("an unrelated identity %v was treated as privileged", other) + } +} diff --git a/client/internal/ipcauth/self_unix.go b/client/internal/ipcauth/self_unix.go new file mode 100644 index 000000000..1b86c4fc0 --- /dev/null +++ b/client/internal/ipcauth/self_unix.go @@ -0,0 +1,17 @@ +//go:build !windows + +package ipcauth + +import "os" + +// CurrentProcessIdentity returns this process's identity as the daemon would +// see it if this process connected to the local IPC. It lets a client (the UI) +// decide up front whether a privileged operation can succeed, without a +// round-trip and without duplicating the rules: the answer comes from the same +// Identity.IsPrivileged the daemon applies. +func CurrentProcessIdentity() (Identity, error) { + return Identity{ + UID: uint32(os.Geteuid()), + GID: uint32(os.Getegid()), + }, nil +} diff --git a/client/internal/ipcauth/self_windows.go b/client/internal/ipcauth/self_windows.go new file mode 100644 index 000000000..5474cc101 --- /dev/null +++ b/client/internal/ipcauth/self_windows.go @@ -0,0 +1,35 @@ +//go:build windows + +package ipcauth + +import ( + "fmt" + + "golang.org/x/sys/windows" +) + +// CurrentProcessIdentity returns this process's identity as the daemon would see +// it if this process connected to the local IPC. It lets a client (the UI) +// decide up front whether a privileged operation can succeed, without a +// round-trip and without duplicating the rules: the answer comes from the same +// Identity.IsPrivileged the daemon applies to the token it reads off the pipe. +func CurrentProcessIdentity() (Identity, error) { + // A pseudo-token, so it must not be closed. + token := windows.GetCurrentProcessToken() + + user, err := token.GetTokenUser() + if err != nil { + return Identity{}, fmt.Errorf("read token user: %w", err) + } + + groups, err := tokenGroupSIDs(token) + if err != nil { + return Identity{}, err + } + + return Identity{ + SID: user.User.Sid.String(), + Groups: groups, + Elevated: token.IsElevated(), + }, nil +} diff --git a/client/internal/profilemanager/config.go b/client/internal/profilemanager/config.go index a110e4102..e1668238e 100644 --- a/client/internal/profilemanager/config.go +++ b/client/internal/profilemanager/config.go @@ -746,6 +746,13 @@ func (config *Config) applyMDMPolicy(policy *mdm.Policy) { // appended for https or ":80" for http. The serviceName parameter is // used to contextualise error messages. On success returns the parsed // *url.URL; on failure returns a non-nil error. +// ParseServiceURL normalises a service URL exactly as the config layer does when +// it stores one, so callers comparing a requested URL against a stored one do not +// have to reimplement the scheme validation and default-port handling. +func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) { + return parseURL(serviceName, serviceURL) +} + func parseURL(serviceName, serviceURL string) (*url.URL, error) { parsedMgmtURL, err := url.ParseRequestURI(serviceURL) if err != nil { diff --git a/client/server/login_gate_test.go b/client/server/login_gate_test.go new file mode 100644 index 000000000..de62a8180 --- /dev/null +++ b/client/server/login_gate_test.go @@ -0,0 +1,127 @@ +package server + +import ( + "context" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/proto" +) + +// A refused login must not leave the profile switched. Login can both switch +// profiles and carry the guarded config fields, so the gate has to run before the +// switch: otherwise a caller whose change is refused still gets the side effect of +// activating whichever profile the request named. +func TestLogin_RefusedChangeLeavesTheProfileAlone(t *testing.T) { + s, _, activeProfile, username, _ := setupServerWithProfile(t) + + // Login reads process state off the daemon's root context. + s.rootCtx = internal.CtxInitState(context.Background()) + + // A second profile that runs the SSH server, which is what makes repointing + // its management binding a privileged change. + target := "ssh-enabled" + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: filepath.Join(profilemanager.DefaultConfigPathDir, target+".json"), + ManagementURL: "https://api.netbird.io:443", + ServerSSHAllowed: boolPtr(true), + }) + require.NoError(t, err) + + _, err = s.Login(userCtx(), &proto.LoginRequest{ + ProfileName: &target, + Username: &username, + ManagementUrl: "https://mgmt.attacker.example:443", + }) + require.Error(t, err, "an unprivileged caller must not move the management URL of an SSH-enabled profile") + require.Equal(t, codes.PermissionDenied, gstatus.Code(err), "want a privilege refusal, got %v", err) + + active, err := s.profileManager.GetActiveProfileState() + require.NoError(t, err) + require.Equal(t, profilemanager.ID(activeProfile), active.ID, + "the refused login switched the active profile anyway") +} + +// A caller whose change becomes privileged only after its first check must be +// refused without having cancelled a login or switched profiles: the first check is +// unsynchronized, so the SSH server can be enabled by a concurrent privileged +// request in between, and the authoritative check happens before any side effect. +func TestLogin_ChangeThatBecomesPrivilegedMidRequestHasNoSideEffects(t *testing.T) { + s, _, activeProfile, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + // The target profile has SSH off, so the first check lets the request through. + target := "ssh-later" + targetPath := filepath.Join(profilemanager.DefaultConfigPathDir, target+".json") + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: targetPath, + ManagementURL: "https://api.netbird.io:443", + ServerSSHAllowed: boolPtr(false), + }) + require.NoError(t, err) + + cancelled := false + s.actCancel = func() { cancelled = true } + + // Stand in for a privileged SetConfig that enables the SSH server between the + // two checks, which is the interleaving the lock has to make safe. + afterLoginPreCheck = func() { + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: targetPath, + ServerSSHAllowed: boolPtr(true), + }) + require.NoError(t, err) + } + t.Cleanup(func() { afterLoginPreCheck = nil }) + + _, err = s.Login(userCtx(), &proto.LoginRequest{ + ProfileName: &target, + Username: &username, + ManagementUrl: "https://mgmt.attacker.example:443", + }) + require.Error(t, err) + require.Equal(t, codes.PermissionDenied, gstatus.Code(err), "want a privilege refusal, got %v", err) + require.False(t, cancelled, "the refused login cancelled the login already in progress") + + active, err := s.profileManager.GetActiveProfileState() + require.NoError(t, err) + require.Equal(t, profilemanager.ID(activeProfile), active.ID, "the refused login switched the active profile anyway") + + stored, err := profilemanager.ReadConfig(targetPath) + require.NoError(t, err) + require.Equal(t, "https://api.netbird.io:443", stored.ManagementURL.String(), "the refused login moved the management URL") +} + +// Login cancels whatever login is already in progress before starting its own. A +// refused caller must not get that far, otherwise anyone able to reach the socket +// can abort someone else's login by sending a request that is denied. +func TestLogin_RefusedChangeLeavesAnInProgressLoginAlone(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + target := "ssh-enabled" + _, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{ + ConfigPath: filepath.Join(profilemanager.DefaultConfigPathDir, target+".json"), + ManagementURL: "https://api.netbird.io:443", + ServerSSHAllowed: boolPtr(true), + }) + require.NoError(t, err) + + cancelled := false + s.actCancel = func() { cancelled = true } + + _, err = s.Login(userCtx(), &proto.LoginRequest{ + ProfileName: &target, + Username: &username, + ManagementUrl: "https://mgmt.attacker.example:443", + }) + require.Error(t, err) + require.Equal(t, codes.PermissionDenied, gstatus.Code(err), "want a privilege refusal, got %v", err) + require.False(t, cancelled, "the refused login cancelled the login already in progress") +} diff --git a/client/server/server.go b/client/server/server.go index 8047006fe..db5909272 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -82,6 +82,12 @@ type Server struct { // extend flow or vice versa. extendAuthSessionFlow *auth.PendingFlow + // guardedConfigMu serializes a privilege check against the write it + // authorizes. Without it the two are separate steps over the same file, and a + // change that was allowed because the profile had the SSH server disabled + // could land after a concurrent privileged request enabled it. + guardedConfigMu sync.Mutex + mutex sync.Mutex config *profilemanager.Config proto.UnimplementedDaemonServiceServer @@ -411,6 +417,20 @@ func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigReques return nil, err } + // Privilege gate: refuse the parts of the request that would let a local + // user turn the root daemon into a root shell. Held across the write so the + // config cannot gain the SSH server between the decision and the update. + s.guardedConfigMu.Lock() + defer s.guardedConfigMu.Unlock() + + stored, err := s.storedProfileConfig(msg.ProfileName, msg.Username) + if err != nil { + return nil, err + } + if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromSetConfig(msg)); err != nil { + return nil, err + } + config, err := s.setConfigInputFromRequest(msg) if err != nil { return nil, err @@ -537,22 +557,23 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro } } - s.mutex.Lock() - if s.actCancel != nil { - s.actCancel() - } - ctx, cancel := context.WithCancel(callerCtx) - - md, ok := metadata.FromIncomingContext(callerCtx) - if ok { - ctx = metadata.NewOutgoingContext(ctx, md) + activeProf, err := s.profileManager.GetActiveProfileState() + if err != nil { + log.Errorf("failed to get active profile state: %v", err) + return nil, fmt.Errorf("failed to get active profile state: %w", err) } - s.actCancel = cancel - s.mutex.Unlock() - - if err := RestoreResidualState(s.rootCtx, s.profileManager.GetStatePath()); err != nil { - log.Warnf(errRestoreResidualState, err) + // Privilege gate: same restrictions as SetConfig, since LoginRequest can carry + // the same fields. It runs before anything here changes daemon state, so a + // refused login neither switches the profile nor cancels a login already in + // progress, and it reads the profile the request targets, which is the one the + // switch below would activate. + stored, err := s.storedLoginConfig(activeProf, msg) + if err != nil { + return nil, err + } + if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil { + return nil, err } state := internal.CtxGetState(s.rootCtx) @@ -563,23 +584,16 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro } }() - activeProf, err := s.profileManager.GetActiveProfileState() + ctx, activeProf, err := s.authorizeAndPrepareLogin(callerCtx, msg, activeProf) if err != nil { - log.Errorf("failed to get active profile state: %v", err) - return nil, fmt.Errorf("failed to get active profile state: %w", err) - } - - if msg.ProfileName != nil { - if _, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf); err != nil { - log.Errorf("failed to switch profile: %v", err) - return nil, err + // The RPC boundary is where this gets recorded: nothing logs handler + // errors for us, and a caller that retries would otherwise leave no + // trace in the daemon log. A refusal is skipped because the gate has + // already logged the decision, with the caller's identity. + if gstatus.Code(err) != codes.PermissionDenied { + log.Errorf("failed to prepare login: %v", err) } - } - - activeProf, err = s.profileManager.GetActiveProfileState() - if err != nil { - log.Errorf("failed to get active profile state: %v", err) - return nil, fmt.Errorf("failed to get active profile state: %w", err) + return nil, err } log.Infof("active profile: %s for %s", activeProf.ID, activeProf.Username) @@ -593,11 +607,6 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro s.mutex.Unlock() - if err := persistLoginOverrides(activeProf, msg.ManagementUrl, msg.OptionalPreSharedKey); err != nil { - log.Errorf("failed to persist login overrides: %v", err) - return nil, fmt.Errorf("persist login overrides: %w", err) - } - config, _, err := s.getConfig(activeProf) if err != nil { log.Errorf("failed to get active profile config: %v", err) @@ -980,6 +989,63 @@ func (s *Server) waitForUp(callerCtx context.Context) (*proto.UpResponse, error) } } +// storedProfileConfig loads the on-disk config of the profile a request +// targets, so a privileged-change decision can be made against the values the +// profile currently holds. A profile that has no config file yet yields nil, +// which every caller must read as "nothing enabled yet". +func (s *Server) storedProfileConfig(handle, username string) (*profilemanager.Config, error) { + resolved, err := s.resolveProfileHandle(handle, username) + if err != nil { + return nil, err + } + + path := resolved.Path + if path == "" { + path = profilemanager.DefaultConfigPath + } + + return s.storedConfigAtPath(path) +} + +// storedLoginConfig loads the on-disk config of the profile a login request +// targets: the one it names, or the active one when it names none. Used to decide +// a privileged change before the request is allowed to switch profiles. +func (s *Server) storedLoginConfig(activeProf *profilemanager.ActiveProfileState, msg *proto.LoginRequest) (*profilemanager.Config, error) { + if msg.ProfileName == nil { + cfgPath, err := activeProf.FilePath() + if err != nil { + return nil, fmt.Errorf("active profile file path: %w", err) + } + return s.storedConfigAtPath(cfgPath) + } + + // Mirrors switchProfileIfNeeded: the default profile resolves without a + // username, so this reads the same profile the switch would activate. + handle := *msg.ProfileName + username := "" + if handle != profilemanager.DefaultProfileName { + username = msg.GetUsername() + } + return s.storedProfileConfig(handle, username) +} + +// storedConfigAtPath reads a profile config file, yielding nil when it does not +// exist yet. +func (s *Server) storedConfigAtPath(path string) (*profilemanager.Config, error) { + if _, err := os.Stat(path); err != nil { + if os.IsNotExist(err) { + return nil, nil //nolint:nilnil + } + return nil, fmt.Errorf("stat profile config: %w", err) + } + + cfg, err := profilemanager.GetConfig(path) + if err != nil { + return nil, fmt.Errorf("read profile config: %w", err) + } + return cfg, nil +} + // resolveProfileHandle resolves a wire-level profile handle (display // name, ID, or unique ID prefix) to a concrete profile. Returns gRPC // status errors so handlers can return them directly. @@ -1197,6 +1263,12 @@ func (s *Server) handleProfileLogout(ctx context.Context, msg *proto.LogoutReque if err := s.logoutFromProfile(ctx, resolved); err != nil { log.Errorf("failed to logout from profile %s: %v", resolved.ID, err) + // A refused deregistration is already a status error carrying the reason + // and the command to run; rewrapping it as Internal would flatten both + // into a gRPC dump for the user. + if _, isStatus := gstatus.FromError(err); isStatus { + return nil, err + } return nil, gstatus.Errorf(codes.Internal, "logout: %v", err) } @@ -1318,6 +1390,13 @@ func (s *Server) sendLogoutRequest(ctx context.Context) error { } func (s *Server) sendLogoutRequestWithConfig(ctx context.Context, config *profilemanager.Config) error { + // Privilege gate: deregistering frees this machine's key to be registered + // against another management server, which is only restricted while the SSH + // server makes that a privilege handover. + if err := requirePrivilegeForDeregistration(ctx, config); err != nil { + return err + } + key, err := wgtypes.ParseKey(config.PrivateKey) if err != nil { return fmt.Errorf("parse private key: %w", err) @@ -2063,7 +2142,10 @@ func (s *Server) RemoveProfile(ctx context.Context, msg *proto.RemoveProfileRequ } if err := s.logoutFromProfile(ctx, resolved); err != nil { - log.Warnf("failed to logout from profile %s before removal: %v", resolved.ID, err) + // Deregistration is best-effort here: the local profile is removed + // either way, so an unprivileged caller leaves the peer registered on + // the management server rather than being blocked from removing it. + log.Warnf("removing profile %s locally without deregistering it: %v", resolved.ID, err) } if err := s.profileManager.RemoveProfile(resolved.ID, msg.Username); err != nil { @@ -2360,6 +2442,69 @@ func sendTerminalNotification() error { // persistLoginOverrides writes management URL and pre-shared key from a LoginRequest to the // active profile config so that subsequent reads pick them up. Empty/nil values are ignored. +// afterLoginPreCheck is a seam for tests to run a concurrent config change +// between Login's first privilege check and the authoritative one. +var afterLoginPreCheck func() + +// authorizeAndPrepareLogin makes the authoritative privilege decision for a login +// and, when it passes, carries out every state change that decision authorizes: +// cancelling an login already in progress, switching to the requested profile, and +// persisting the config overrides the request carries. +// +// All of it happens under guardedConfigMu, which SetConfig also holds across its +// own check and write. Login's earlier check refuses the ordinary case before any +// of this is reached; this one exists because that check is not synchronized +// against a concurrent privileged request that enables the SSH server, and a +// caller refused here must not have cancelled or switched anything either. +func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto.LoginRequest, activeProf *profilemanager.ActiveProfileState) (context.Context, *profilemanager.ActiveProfileState, error) { + if afterLoginPreCheck != nil { + afterLoginPreCheck() + } + + s.guardedConfigMu.Lock() + defer s.guardedConfigMu.Unlock() + + stored, err := s.storedLoginConfig(activeProf, msg) + if err != nil { + return nil, nil, err + } + if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil { + return nil, nil, err + } + + s.mutex.Lock() + if s.actCancel != nil { + s.actCancel() + } + ctx, cancel := context.WithCancel(callerCtx) + if md, ok := metadata.FromIncomingContext(callerCtx); ok { + ctx = metadata.NewOutgoingContext(ctx, md) + } + s.actCancel = cancel + s.mutex.Unlock() + + if err := RestoreResidualState(s.rootCtx, s.profileManager.GetStatePath()); err != nil { + log.Warnf(errRestoreResidualState, err) + } + + if msg.ProfileName != nil { + if _, err := s.switchProfileIfNeeded(*msg.ProfileName, msg.Username, activeProf); err != nil { + return nil, nil, fmt.Errorf("switch profile: %w", err) + } + } + + activeProf, err = s.profileManager.GetActiveProfileState() + if err != nil { + return nil, nil, fmt.Errorf("active profile state: %w", err) + } + + if err := persistLoginOverrides(activeProf, msg.ManagementUrl, msg.OptionalPreSharedKey); err != nil { + return nil, nil, fmt.Errorf("persist login overrides: %w", err) + } + + return ctx, activeProf, nil +} + func persistLoginOverrides(activeProf *profilemanager.ActiveProfileState, managementURL string, preSharedKey *string) error { if preSharedKey != nil && *preSharedKey == "" { preSharedKey = nil diff --git a/client/server/setconfig_mdm_test.go b/client/server/setconfig_mdm_test.go index 9baf16136..ae323ea8c 100644 --- a/client/server/setconfig_mdm_test.go +++ b/client/server/setconfig_mdm_test.go @@ -66,7 +66,11 @@ func setupServerWithProfile(t *testing.T) (s *Server, ctx context.Context, profN Username: currUser.Username, })) - ctx = context.Background() + // The privileged-change gate reads the caller's kernel identity from the + // context, which a real caller gets from the daemon's transport credentials. + // This test drives the handler directly, so it stands in for a root caller; + // without an identity the gate would (correctly) refuse the SSH fields. + ctx = privilegedTestCtx() s = New(ctx, "console", "", false, false, false, false) return s, ctx, profName, currUser.Username, cfgPath } diff --git a/client/server/setconfig_test.go b/client/server/setconfig_test.go index 0e55257a9..db7a26f03 100644 --- a/client/server/setconfig_test.go +++ b/client/server/setconfig_test.go @@ -1,7 +1,6 @@ package server import ( - "context" "os/user" "path/filepath" "reflect" @@ -52,7 +51,11 @@ func TestSetConfig_AllFieldsSaved(t *testing.T) { }) require.NoError(t, err) - ctx := context.Background() + // The privileged-change gate reads the caller's kernel identity from the + // context, which a real caller gets from the daemon's transport credentials. + // This test drives the handler directly, so it stands in for a root caller; + // without an identity the gate would (correctly) refuse the SSH fields. + ctx := privilegedTestCtx() s := New(ctx, "console", "", false, false, false, false) rosenpassEnabled := true diff --git a/client/server/ssh_gate.go b/client/server/ssh_gate.go new file mode 100644 index 000000000..ca1b4c4ee --- /dev/null +++ b/client/server/ssh_gate.go @@ -0,0 +1,282 @@ +package server + +import ( + "context" + "fmt" + "net/url" + "runtime" + "strings" + + log "github.com/sirupsen/logrus" + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/client/proto" + "github.com/netbirdio/netbird/util" +) + +// The daemon runs as root/LocalSystem, so a handful of config changes cross the +// user-to-root boundary and are restricted to privileged callers: +// +// - Enabling SSH root login, or disabling SSH authentication, turns the +// daemon's SSH server into a root (or unauthenticated) shell. +// - Enabling the SSH server at all is what makes the above reachable, and a +// profile the caller owns is not a privilege they hold. +// - While the SSH server is enabled, repointing the profile at another +// management identity hands SSH authorization decisions, including which +// keys and users are accepted, to whoever controls that identity. Changing +// the management URL and deregistering the peer are both ways to do that. +// +// Everything else stays unauthenticated, so this is not an authorization model: +// it only refuses the changes that would let a local user become root. A caller +// whose identity cannot be established is refused as well. + +// privilegedConfigChange is the subset of a config request that crosses the +// user-to-root boundary. Fields are nil or empty when the request leaves them +// untouched. +type privilegedConfigChange struct { + managementURL string + serverSSHAllowed *bool + enableSSHRoot *bool + disableSSHAuth *bool +} + +func privilegedChangeFromSetConfig(msg *proto.SetConfigRequest) privilegedConfigChange { + return privilegedConfigChange{ + managementURL: msg.GetManagementUrl(), + serverSSHAllowed: msg.ServerSSHAllowed, + enableSSHRoot: msg.EnableSSHRoot, + disableSSHAuth: msg.DisableSSHAuth, + } +} + +func privilegedChangeFromLogin(msg *proto.LoginRequest) privilegedConfigChange { + return privilegedConfigChange{ + managementURL: msg.GetManagementUrl(), + serverSSHAllowed: msg.ServerSSHAllowed, + enableSSHRoot: msg.EnableSSHRoot, + disableSSHAuth: msg.DisableSSHAuth, + } +} + +// requirePrivilegeForConfigChange refuses the privileged parts of a config +// change when the caller is not root/administrator. stored is the profile's +// current config, or nil when it has none yet. +// +// Each check compares against the stored value so that a request restating a +// value it does not change is never refused: a UI that submits the whole +// settings form must not start failing once an administrator has enabled SSH. +func requirePrivilegeForConfigChange(ctx context.Context, stored *profilemanager.Config, change privilegedConfigChange) error { + if enables(storedFlag(stored, func(c *profilemanager.Config) *bool { return c.EnableSSHRoot }), change.enableSSHRoot) { + return denyPrivileged(ctx, "enabling SSH root login", ipcauth.UpCommand("--enable-ssh-root")) + } + + if enables(storedFlag(stored, func(c *profilemanager.Config) *bool { return c.DisableSSHAuth }), change.disableSSHAuth) { + return denyPrivileged(ctx, "disabling SSH authentication", ipcauth.UpCommand("--disable-ssh-auth")) + } + + if enables(sshServerCurrentlyAllowed(stored), change.serverSSHAllowed) { + return denyPrivileged(ctx, "enabling the NetBird SSH server", ipcauth.UpCommand("--allow-server-ssh")) + } + + // Only guard the management binding while the SSH server is enabled: that is + // when the management identity decides who may open a shell here. + if !sshServerEnabled(stored) { + return nil + } + + if change.managementURL != "" && !sameManagementURL(stored.ManagementURL, change.managementURL) { + return denyPrivileged(ctx, + "changing the management URL while the NetBird SSH server is enabled", + ipcauth.UpCommand("-m "+change.managementURL)) + } + + return nil +} + +// requirePrivilegeForDeregistration refuses to deregister the peer from the +// management server when the caller is not privileged and the profile has the +// SSH server enabled. Deregistering frees the peer's key to be registered +// against another management identity, which is the same handover the +// management URL check refuses. +// +// Callers that treat deregistration as best-effort (profile removal) continue +// without it; callers that were asked to deregister surface the error. +func requirePrivilegeForDeregistration(ctx context.Context, cfg *profilemanager.Config) error { + if !sshServerEnabled(cfg) { + return nil + } + + return denyPrivileged(ctx, + "deregistering this peer while the NetBird SSH server is enabled", + ipcauth.ElevatedCommand("netbird logout")) +} + +// denyPrivileged returns nil when the caller is privileged, and otherwise a +// PermissionDenied whose message names the action and the command that performs +// it with the privileges it needs. The same summary and command ride along as an +// ErrorInfo detail so the CLI and the UI can present them without parsing text. +// +// action reads as the subject of a sentence ("enabling SSH root login"), and +// command is the equivalent command, already elevated for the platform. +func denyPrivileged(ctx context.Context, action, command string) error { + id, ok := ipcauth.CallerIdentity(ctx) + if !ok { + log.Warnf("denying %s: the caller's identity cannot be verified on this control channel", action) + return privilegeError(unidentifiedSummary(action), reinstallCommand()) + } + + if ipcauth.IsPrivilegedCaller(id) { + log.Infof("allowing %s for privileged caller %s", action, id) + return nil + } + + log.Warnf("denying %s for unprivileged caller %s", action, id) + actor, command := requiredActor(command) + return privilegeError(privilegeSummary(action, actor), command) +} + +// requiredActor names who may perform the operation and adjusts the command to +// match. A daemon that is not itself privileged delegates to its own identity, so +// telling that host's user to become root is wrong twice over: root is not what the +// daemon checks for, and a rootless container has neither root nor sudo. +func requiredActor(command string) (string, string) { + self, delegates := ipcauth.SelfDelegatesTo() + if !delegates { + return ipcauth.PrivilegedActor(), command + } + return fmt.Sprintf("the user the daemon runs as (%s)", self), strings.ReplaceAll(command, "sudo ", "") +} + +// privilegeError builds the PermissionDenied carrying summary and command. +func privilegeError(summary, command string) error { + st := gstatus.New(codes.PermissionDenied, fmt.Sprintf("%s\n\n%s", summary, command)) + + detailed, err := st.WithDetails(&errdetails.ErrorInfo{ + Reason: ipcauth.ErrorReasonPrivilegeRequired, + Domain: ipcauth.ErrorDomain, + Metadata: map[string]string{ + ipcauth.ErrorMetaSummary: summary, + ipcauth.ErrorMetaCommand: command, + }, + }) + if err != nil { + log.Debugf("attach privilege error detail: %v", err) + return st.Err() + } + return detailed.Err() +} + +// privilegeSummary states what is refused and what it needs, in one sentence +// that reads the same in a dialog and in a terminal. +func privilegeSummary(action, actor string) string { + return fmt.Sprintf("%s requires %s.", capitalize(action), actor) +} + +// unidentifiedSummary covers a control channel that carries no caller identity. +// Elevating does not help there, so it points at the daemon's socket instead. +func unidentifiedSummary(action string) string { + return fmt.Sprintf("%s requires %s, and the daemon cannot verify who is calling over its current socket. "+ + "Reinstall the service on a socket that carries the caller's identity.", capitalize(action), ipcauth.PrivilegedActor()) +} + +// reinstallCommand is the command that moves the daemon onto a socket whose +// callers can be identified. +func reinstallCommand() string { + if runtime.GOOS == "windows" { + return fmt.Sprintf("netbird service install --daemon-addr %s", daemonaddr.WindowsPipeAddr) + } + return "sudo netbird service install --daemon-addr unix:///var/run/netbird.sock" +} + +func capitalize(s string) string { + if s == "" { + return s + } + return strings.ToUpper(s[:1]) + s[1:] +} + +// enables reports whether requested turns a flag on that is currently off. A +// request that restates the stored value, or turns the flag off, is not a +// privileged change. +func enables(stored, requested *bool) bool { + if requested == nil || !*requested { + return false + } + return stored == nil || !*stored +} + +// storedFlag reads a flag from the stored config, tolerating a config that does +// not exist yet. +func storedFlag(cfg *profilemanager.Config, get func(*profilemanager.Config) *bool) *bool { + if cfg == nil { + return nil + } + return get(cfg) +} + +// sshServerEnabled reports whether the profile currently runs the SSH server. +// +// A nil flag means ON, matching what the engine does with the same config +// (util.ReturnBoolWithDefaultTrue in internal/connect.go, kept for configs written +// before the flag existed). Reading it as OFF here would open the management-URL +// and deregistration guards on exactly those legacy hosts, whose SSH server is +// running. Configs loaded through profilemanager have already been materialised by +// apply(), so this is the same answer by a route that does not depend on that. +func sshServerEnabled(cfg *profilemanager.Config) bool { + if cfg == nil { + return false + } + return util.ReturnBoolWithDefaultTrue(cfg.ServerSSHAllowed) +} + +// sshServerCurrentlyAllowed is the value an enable request is compared against. It +// shares sshServerEnabled's nil-means-on default, so restating "on" for a legacy +// config is correctly seen as no change. +func sshServerCurrentlyAllowed(cfg *profilemanager.Config) *bool { + enabled := sshServerEnabled(cfg) + if cfg == nil { + return nil + } + return &enabled +} + +// sameManagementURL reports whether requested addresses the same management +// server as stored, comparing scheme, host and effective port so that an +// equivalent spelling ("https://api.netbird.io" for a stored +// "https://api.netbird.io:443") is not treated as a change. It fails closed: +// anything unparseable counts as a change and therefore needs privilege. +func sameManagementURL(stored *url.URL, requested string) bool { + if stored == nil { + return false + } + + // Normalise the requested URL through the config layer's own parser, so the + // comparison cannot drift from how the value would actually be stored. + parsed, err := profilemanager.ParseServiceURL("Management URL", requested) + if err != nil { + return false + } + + return stored.Scheme == parsed.Scheme && + stored.Hostname() == parsed.Hostname() && + effectivePort(stored) == effectivePort(parsed) +} + +func effectivePort(u *url.URL) string { + if port := u.Port(); port != "" { + return port + } + switch u.Scheme { + case "https": + return "443" + case "http": + return "80" + default: + return "" + } +} diff --git a/client/server/ssh_gate_test.go b/client/server/ssh_gate_test.go new file mode 100644 index 000000000..cbd345f16 --- /dev/null +++ b/client/server/ssh_gate_test.go @@ -0,0 +1,348 @@ +package server + +import ( + "context" + "net/url" + "os" + "runtime" + "strings" + "testing" + + "google.golang.org/genproto/googleapis/rpc/errdetails" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/internal/profilemanager" +) + +// ctxWithIdentity builds a request context carrying the identity the transport +// credentials would have attached. +func ctxWithIdentity(id ipcauth.Identity) context.Context { + return peer.NewContext(context.Background(), &peer.Peer{ + AuthInfo: ipcauth.AuthInfo{ + CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity}, + Identity: id, + }, + }) +} + +// unprivUID is deliberately not this process's own uid. An unprivileged daemon +// treats a caller sharing its identity as privileged (rootless containers), and +// the test binary would otherwise stand in for both the daemon and the caller. +// os.Geteuid returns -1 on Windows, where identities are SIDs instead and this is +// unused. +var unprivUID = uint32(os.Geteuid() + 1) + +// The fabricated identities have to be shaped like the platform's: a uid says +// nothing on Windows, and a zero uid there would read as root and be privileged. +func rootCtx() context.Context { return ctxWithIdentity(privilegedIdentity()) } +func userCtx() context.Context { return ctxWithIdentity(unprivilegedIdentity()) } + +func privilegedIdentity() ipcauth.Identity { + if runtime.GOOS == "windows" { + // LocalSystem, which is what the Windows service account is. + return ipcauth.Identity{SID: "S-1-5-18"} + } + return ipcauth.Identity{UID: 0} +} + +func unprivilegedIdentity() ipcauth.Identity { + if runtime.GOOS == "windows" { + // A plain user SID: no groups, so no BUILTIN\Administrators, and not + // elevated. + return ipcauth.Identity{SID: "S-1-5-21-1-2-3-1001"} + } + return ipcauth.Identity{UID: unprivUID, GID: unprivUID} +} +func noIdentityCtx() context.Context { return context.Background() } + +func boolPtr(v bool) *bool { return &v } + +func mustURL(t *testing.T, raw string) *url.URL { + t.Helper() + u, err := url.Parse(raw) + if err != nil { + t.Fatalf("parse %q: %v", raw, err) + } + return u +} + +func assertDenied(t *testing.T, err error) { + t.Helper() + if err == nil { + t.Fatal("expected the change to be refused, got nil") + } + st := gstatus.Convert(err) + if st.Code() != codes.PermissionDenied { + t.Fatalf("code = %v, want PermissionDenied", st.Code()) + } + // The refusal must be machine-readable: the CLI and the UI render the + // summary and command from the detail rather than parsing the message. + var info *errdetails.ErrorInfo + for _, d := range st.Details() { + if got, ok := d.(*errdetails.ErrorInfo); ok { + info = got + } + } + if info == nil { + t.Fatal("refusal carries no ErrorInfo detail") + } + if info.GetReason() != ipcauth.ErrorReasonPrivilegeRequired || info.GetDomain() != ipcauth.ErrorDomain { + t.Fatalf("detail = %s/%s, want %s/%s", info.GetDomain(), info.GetReason(), ipcauth.ErrorDomain, ipcauth.ErrorReasonPrivilegeRequired) + } + if info.GetMetadata()[ipcauth.ErrorMetaSummary] == "" { + t.Error("detail carries no summary") + } + if info.GetMetadata()[ipcauth.ErrorMetaCommand] == "" { + t.Error("detail carries no command") + } +} + +func assertAllowed(t *testing.T, err error) { + t.Helper() + if err != nil { + t.Fatalf("expected the change to be allowed, got %v", err) + } +} + +func TestRequirePrivilegeForConfigChange_SSHFlags(t *testing.T) { + tests := []struct { + name string + stored *profilemanager.Config + change privilegedConfigChange + privileged bool + wantDeny bool + }{ + { + name: "enabling the ssh server unprivileged is refused", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + wantDeny: true, + }, + { + name: "enabling the ssh server as root is allowed", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + privileged: true, + }, + { + name: "restating an already enabled ssh server is not a change", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + }, + { + name: "turning the ssh server off is not guarded", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(false)}, + }, + { + name: "a profile with no config yet counts as off, so enabling is refused", + stored: nil, + change: privilegedConfigChange{serverSSHAllowed: boolPtr(true)}, + wantDeny: true, + }, + { + name: "enabling ssh root login unprivileged is refused", + stored: &profilemanager.Config{EnableSSHRoot: boolPtr(false)}, + change: privilegedConfigChange{enableSSHRoot: boolPtr(true)}, + wantDeny: true, + }, + { + name: "restating ssh root login is not a change", + stored: &profilemanager.Config{EnableSSHRoot: boolPtr(true)}, + change: privilegedConfigChange{enableSSHRoot: boolPtr(true)}, + }, + { + name: "turning ssh root login off is not guarded", + stored: &profilemanager.Config{EnableSSHRoot: boolPtr(true)}, + change: privilegedConfigChange{enableSSHRoot: boolPtr(false)}, + }, + { + name: "disabling ssh authentication unprivileged is refused", + stored: &profilemanager.Config{DisableSSHAuth: boolPtr(false)}, + change: privilegedConfigChange{disableSSHAuth: boolPtr(true)}, + wantDeny: true, + }, + { + name: "re-enabling ssh authentication is not guarded", + stored: &profilemanager.Config{DisableSSHAuth: boolPtr(true)}, + change: privilegedConfigChange{disableSSHAuth: boolPtr(false)}, + }, + { + name: "a request that touches none of the guarded fields is allowed", + stored: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + change: privilegedConfigChange{}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := userCtx() + if tt.privileged { + ctx = rootCtx() + } + err := requirePrivilegeForConfigChange(ctx, tt.stored, tt.change) + if tt.wantDeny { + assertDenied(t, err) + return + } + assertAllowed(t, err) + }) + } +} + +func TestRequirePrivilegeForConfigChange_ManagementURL(t *testing.T) { + sshOn := func(raw string) *profilemanager.Config { + return &profilemanager.Config{ServerSSHAllowed: boolPtr(true), ManagementURL: mustURL(t, raw)} + } + sshOff := func(raw string) *profilemanager.Config { + return &profilemanager.Config{ServerSSHAllowed: boolPtr(false), ManagementURL: mustURL(t, raw)} + } + + tests := []struct { + name string + stored *profilemanager.Config + requested string + privileged bool + wantDeny bool + }{ + { + name: "moving the binding while ssh is enabled is refused", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://attacker.example.com:443", + wantDeny: true, + }, + { + name: "moving the binding as root is allowed", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://selfhosted.example.com:443", + privileged: true, + }, + { + name: "the same url restated is not a change", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://api.netbird.io:443", + }, + { + name: "an equivalent spelling of the same url is not a change", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://api.netbird.io", + }, + { + name: "an equivalent spelling with an explicit http port is not a change", + stored: sshOn("http://mgmt.internal:80"), + requested: "http://mgmt.internal", + }, + { + name: "a different port on the same host is a change", + stored: sshOn("https://api.netbird.io:443"), + requested: "https://api.netbird.io:8443", + wantDeny: true, + }, + { + name: "a different scheme on the same host is a change", + stored: sshOn("https://mgmt.internal:443"), + requested: "http://mgmt.internal:443", + wantDeny: true, + }, + { + name: "with ssh disabled the binding is not guarded at all", + stored: sshOff("https://api.netbird.io:443"), + requested: "https://attacker.example.com:443", + }, + { + name: "an unparseable url fails closed", + stored: sshOn("https://api.netbird.io:443"), + requested: "ht tp://%zz", + wantDeny: true, + }, + { + name: "an empty url leaves the binding alone", + stored: sshOn("https://api.netbird.io:443"), + requested: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := userCtx() + if tt.privileged { + ctx = rootCtx() + } + err := requirePrivilegeForConfigChange(ctx, tt.stored, privilegedConfigChange{managementURL: tt.requested}) + if tt.wantDeny { + assertDenied(t, err) + return + } + assertAllowed(t, err) + }) + } +} + +// A caller the daemon cannot identify must be refused, not trusted: that is the +// state on a TCP daemon socket, where no peer credentials exist. +func TestRequirePrivilegeForConfigChange_UnidentifiedCallerIsRefused(t *testing.T) { + err := requirePrivilegeForConfigChange(noIdentityCtx(), + &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + privilegedConfigChange{serverSSHAllowed: boolPtr(true)}) + assertDenied(t, err) + + // The guidance must point at the socket rather than at sudo, since elevating + // would not help. + st := gstatus.Convert(err) + if !strings.Contains(st.Message(), "service install") { + t.Errorf("message %q does not tell the operator how to fix the socket", st.Message()) + } +} + +func TestRequirePrivilegeForDeregistration(t *testing.T) { + tests := []struct { + name string + cfg *profilemanager.Config + privileged bool + wantDeny bool + }{ + { + name: "deregistering while ssh is enabled is refused", + cfg: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + wantDeny: true, + }, + { + name: "deregistering while ssh is enabled is allowed for root", + cfg: &profilemanager.Config{ServerSSHAllowed: boolPtr(true)}, + privileged: true, + }, + { + name: "deregistering with ssh disabled is not guarded", + cfg: &profilemanager.Config{ServerSSHAllowed: boolPtr(false)}, + }, + { + name: "deregistering a profile with no config is not guarded", + cfg: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := userCtx() + if tt.privileged { + ctx = rootCtx() + } + err := requirePrivilegeForDeregistration(ctx, tt.cfg) + if tt.wantDeny { + assertDenied(t, err) + return + } + assertAllowed(t, err) + }) + } +} + +// privilegedTestCtx is the context a handler-level test should use when it is +// standing in for a root/administrator caller. Tests that drive the handlers +// directly have no transport credentials, and the privileged-change gate refuses +// a caller it cannot identify. +func privilegedTestCtx() context.Context { return rootCtx() } diff --git a/client/ssh/client/client.go b/client/ssh/client/client.go index ebf8eb794..4180849cd 100644 --- a/client/ssh/client/client.go +++ b/client/ssh/client/client.go @@ -9,7 +9,6 @@ import ( "path/filepath" "runtime" "strconv" - "strings" "time" log "github.com/sirupsen/logrus" @@ -17,7 +16,6 @@ import ( "golang.org/x/crypto/ssh/knownhosts" "golang.org/x/term" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/profilemanager" @@ -32,7 +30,7 @@ const ( // DefaultDaemonAddr is the default address for the NetBird daemon DefaultDaemonAddr = "unix:///var/run/netbird.sock" // DefaultDaemonAddrWindows is the default address for the NetBird daemon on Windows - DefaultDaemonAddrWindows = "tcp://127.0.0.1:41731" + DefaultDaemonAddrWindows = daemonaddr.WindowsPipeAddr ) // Client wraps crypto/ssh Client for simplified SSH operations @@ -268,7 +266,7 @@ func getDefaultDaemonAddr() string { return addr } if runtime.GOOS == "windows" { - return DefaultDaemonAddrWindows + return daemonaddr.ResolveDaemonAddr(DefaultDaemonAddrWindows) } return daemonaddr.ResolveUnixDaemonAddr(DefaultDaemonAddr) } @@ -410,12 +408,9 @@ func verifyHostKeyViaDaemon(hostname string, remote net.Addr, key ssh.PublicKey, } func connectToDaemon(daemonAddr string) (*grpc.ClientConn, error) { - addr := strings.TrimPrefix(daemonAddr, "tcp://") + target, opts := daemonaddr.DialTarget(daemonAddr) - conn, err := grpc.NewClient( - addr, - grpc.WithTransportCredentials(insecure.NewCredentials()), - ) + conn, err := grpc.NewClient(target, opts...) if err != nil { log.Debugf("failed to create gRPC client for NetBird daemon at %s: %v", daemonAddr, err) return nil, fmt.Errorf("failed to connect to NetBird daemon: %w", err) diff --git a/client/ssh/proxy/proxy.go b/client/ssh/proxy/proxy.go index 73b50122c..721810edb 100644 --- a/client/ssh/proxy/proxy.go +++ b/client/ssh/proxy/proxy.go @@ -9,7 +9,6 @@ import ( "net" "os" "strconv" - "strings" "sync" "time" @@ -17,8 +16,8 @@ import ( log "github.com/sirupsen/logrus" cryptossh "golang.org/x/crypto/ssh" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/proto" nbssh "github.com/netbirdio/netbird/client/ssh" @@ -55,8 +54,8 @@ type SSHProxy struct { } func New(daemonAddr, targetHost string, targetPort int, stderr io.Writer, browserOpener func(string) error) (*SSHProxy, error) { - grpcAddr := strings.TrimPrefix(daemonAddr, "tcp://") - grpcConn, err := grpc.NewClient(grpcAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) + target, opts := daemonaddr.DialTarget(daemonAddr) + grpcConn, err := grpc.NewClient(target, opts...) if err != nil { return nil, fmt.Errorf("connect to daemon: %w", err) } diff --git a/client/ui/frontend/src/components/CopyToClipboard.tsx b/client/ui/frontend/src/components/CopyToClipboard.tsx index 1b2a87da4..3cf681a1c 100644 --- a/client/ui/frontend/src/components/CopyToClipboard.tsx +++ b/client/ui/frontend/src/components/CopyToClipboard.tsx @@ -18,6 +18,9 @@ type CopyToClipboardProps = { className?: string; iconClassName?: string; alwaysShowIcon?: boolean; + // wrap lets long content (a shell command, a path) break across lines + // instead of being truncated to one line. + wrap?: boolean; variant?: CopyToClipboardVariant; "aria-label"?: string; tabIndex?: number; @@ -32,6 +35,7 @@ export const CopyToClipboard = ({ className, iconClassName, alwaysShowIcon = false, + wrap = false, variant = "default", "aria-label": ariaLabel, tabIndex = 0, @@ -83,7 +87,8 @@ export const CopyToClipboard = ({ > { loadedRef.current = loaded; }, [loaded]); + // reload re-reads the daemon's config, which is authoritative. Used on + // mount, on the daemon's config_changed event, and to undo an optimistic + // update the daemon then rejected. + const reload = useCallback( + async (profileName: string) => { + try { + const data = await SettingsSvc.GetConfig({ profileName, username }); + setLoaded({ profileName, data }); + } catch (e) { + console.warn("[SettingsContext] reload after rejected save failed", e); + } + }, + [username], + ); + useEffect(() => { if (!profileLoaded || !activeProfileId) return; let cancelled = false; @@ -133,13 +148,20 @@ const useSettingsState = () => { username, }); } catch (e) { + // The optimistic update is wrong now: the daemon refused it + // (a change that needs elevated privileges, an MDM-managed + // field, ...). Snap the controls back to what it actually + // holds before reporting, so the UI never shows a value the + // daemon does not have. + await reload(profileName); await errorDialog({ Title: i18next.t("settings.error.saveTitle"), Message: errorMessage(e), + Command: errorCommand(e), }); } }, - [username], + [username, reload], ); const setField = useCallback( diff --git a/client/ui/frontend/src/hooks/usePrivilege.ts b/client/ui/frontend/src/hooks/usePrivilege.ts new file mode 100644 index 000000000..05e9a7ce0 --- /dev/null +++ b/client/ui/frontend/src/hooks/usePrivilege.ts @@ -0,0 +1,32 @@ +import { useEffect, useState } from "react"; +import { Settings as SettingsSvc } from "@bindings/services"; +import { Privilege } from "@bindings/services/models.js"; + +// usePrivilege reports whether this UI process may perform the changes the daemon +// restricts to root/administrator. It is answered in-process from our own token +// with the daemon's own rule, so there is no round-trip and it works while the +// daemon is down. +// +// null means "not known yet", which includes the read having failed. Callers must +// treat that as "do not restrict": the daemon enforces this regardless, so the +// only thing a wrong guess here costs is a control that looks unavailable when it +// is not, or a save that fails with the daemon's own guidance. +export const usePrivilege = (): Privilege | null => { + const [privilege, setPrivilege] = useState(null); + + useEffect(() => { + let cancelled = false; + SettingsSvc.Privilege() + .then((p) => { + if (!cancelled) setPrivilege(p); + }) + .catch((e: unknown) => { + console.warn("[usePrivilege] read failed, not restricting controls", e); + }); + return () => { + cancelled = true; + }; + }, []); + + return privilege; +}; diff --git a/client/ui/frontend/src/lib/errors.ts b/client/ui/frontend/src/lib/errors.ts index 34a90dace..b4dee2717 100644 --- a/client/ui/frontend/src/lib/errors.ts +++ b/client/ui/frontend/src/lib/errors.ts @@ -1,6 +1,6 @@ import { WindowManager } from "@bindings/services"; -type ClassifiedError = { short: string; long: string }; +type ClassifiedError = { short: string; long: string; command: string }; const asObject = (v: unknown): Record | null => v && typeof v === "object" ? (v as Record) : null; @@ -22,20 +22,24 @@ const toWailsEnvelope = (e: unknown): Record | null => { return asObject(obj.cause) ?? parseJsonObject(obj.message); }; -// Read { short, long } from wherever the classified error sits in the envelope +// Read { short, long, command } from wherever the classified error sits in the envelope const toClassifiedError = (v: unknown): ClassifiedError | null => { const o = asObject(v); if (!o) return null; const short = typeof o.short === "string" ? o.short : ""; const long = typeof o.long === "string" ? o.long : ""; - return short || long ? { short, long } : null; + const command = typeof o.command === "string" ? o.command : ""; + return short || long ? { short, long, command } : null; +}; + +const classify = (e: unknown): ClassifiedError | null => { + const envelope = toWailsEnvelope(e); + return toClassifiedError(envelope?.cause) ?? toClassifiedError(envelope); }; export const formatErrorMessage = (e: unknown): string => { - const envelope = toWailsEnvelope(e); - // Prefer the structured { short, long } the daemon classifier produced. - const classified = toClassifiedError(envelope?.cause) ?? toClassifiedError(envelope); + const classified = classify(e); if (classified) { const { short, long } = classified; if (short && long && long !== short) return `${short} Details: ${long}`; @@ -44,17 +48,26 @@ export const formatErrorMessage = (e: unknown): string => { } // Unclassified (a service returned the raw daemon error) + const envelope = toWailsEnvelope(e); const message = envelope?.message; if (typeof message === "string" && message) return message; if (e instanceof Error) return e.message; return String(e); }; +// errorCommand returns a command the user can run to complete an operation the +// daemon refused, when the error carries one (a change that needs elevated +// privileges). Empty for every other error. +export const errorCommand = (e: unknown): string => classify(e)?.command ?? ""; + export type ErrorDialogOptions = { Title: string; Message: string; + // Command is shown for copying below the message. Defaults to the one the + // error carries, so callers only pass it to override. + Command?: string; }; export function errorDialog(options: ErrorDialogOptions): Promise { - return WindowManager.OpenError(options.Title, options.Message); + return WindowManager.OpenError(options.Title, options.Message, options.Command ?? ""); } diff --git a/client/ui/frontend/src/modules/error/ErrorDialog.tsx b/client/ui/frontend/src/modules/error/ErrorDialog.tsx index f7929fad2..4fbb78052 100644 --- a/client/ui/frontend/src/modules/error/ErrorDialog.tsx +++ b/client/ui/frontend/src/modules/error/ErrorDialog.tsx @@ -3,6 +3,7 @@ import { useTranslation } from "react-i18next"; import { useSearchParams } from "react-router-dom"; import { AlertCircleIcon } from "lucide-react"; import { Button } from "@/components/buttons/Button"; +import { CopyToClipboard } from "@/components/CopyToClipboard"; import { ConfirmDialog } from "@/components/dialog/ConfirmDialog"; import { DialogActions } from "@/components/dialog/DialogActions"; import { DialogDescription } from "@/components/dialog/DialogDescription"; @@ -12,14 +13,22 @@ import { WindowManager } from "@bindings/services"; import { useAutoSizeWindow } from "@/hooks/useAutoSizeWindow"; const WINDOW_WIDTH = 380; +// A command needs the room to wrap at a sensible number of characters instead of +// breaking every few words. +const WINDOW_WIDTH_WITH_COMMAND = 460; export default function ErrorDialog() { const { t } = useTranslation(); - const contentRef = useAutoSizeWindow(WINDOW_WIDTH); const [params] = useSearchParams(); const title = params.get("title") || t("window.title.error"); const message = params.get("message") || ""; + // Set when the daemon refused an operation that needs elevated privileges: + // the command that performs it, offered for copying. + const command = params.get("command") || ""; + const contentRef = useAutoSizeWindow( + command ? WINDOW_WIDTH_WITH_COMMAND : WINDOW_WIDTH, + ); const close = useCallback(() => { WindowManager.CloseError().catch(console.error); @@ -37,15 +46,37 @@ export default function ErrorDialog() { -
+
{title} {message && ( - {message} + {/* select-text: the message often names a path, a flag or an + address the user needs to act on. */} + + {message} + )} + {command && ( + + + {command} + + + )}
diff --git a/client/ui/frontend/src/modules/settings/SettingsSSH.tsx b/client/ui/frontend/src/modules/settings/SettingsSSH.tsx index f18fd9493..bd91e520c 100644 --- a/client/ui/frontend/src/modules/settings/SettingsSSH.tsx +++ b/client/ui/frontend/src/modules/settings/SettingsSSH.tsx @@ -1,4 +1,5 @@ import { useTranslation } from "react-i18next"; +import { CopyToClipboard } from "@/components/CopyToClipboard"; import FancyToggleSwitch from "@/components/switches/FancyToggleSwitch"; import { HelpText } from "@/components/typography/HelpText"; import { Input } from "@/components/inputs/Input"; @@ -6,12 +7,50 @@ import { Label } from "@/components/typography/Label"; import { cn } from "@/lib/cn"; import { SectionGroup } from "@/modules/settings/SettingsSection.tsx"; import { useSettings } from "@/contexts/SettingsContext.tsx"; -import { type ChangeEvent, useEffect, useId, useState } from "react"; +import { usePrivilege } from "@/hooks/usePrivilege.ts"; +import { Privilege } from "@bindings/services/models.js"; +import { type ChangeEvent, type ReactNode, useEffect, useId, useState } from "react"; export function SettingsSSH() { const { t } = useTranslation(); const { config, setField } = useSettings(); + const privilege = usePrivilege(); const isSSHServerEnabled = config.serverSshAllowed; + + // The daemon restricts only the direction that hands out shells from a process + // running as root. So for an unprivileged user a guarded control is either + // unavailable (it is off and only they could turn it on) or a one-way switch + // (it is on, they may turn it off, but not back on) — say which, either way. + // + // A null privilege means we could not determine it: leave the control alone + // rather than greying it out with nothing to explain why. The daemon enforces + // this regardless, and a rejected save reports its own guidance. + const guarded = ( + guardedDirectionActive: boolean, + command: (p: Privilege) => string, + // inverted marks a control whose guarded direction is switching it off, so + // the one-way warning has to read the other way round. + inverted = false, + ) => { + if (!privilege || privilege.privileged) { + return { disabled: false, hint: undefined }; + } + const hint = ( + + ); + return { disabled: !guardedDirectionActive, hint }; + }; + + const sshServer = guarded(config.serverSshAllowed, (p) => p.allowSshServer); + const sshRoot = guarded(config.enableSshRoot, (p) => p.enableSshRoot); + // Inverted control: the guarded direction is switching authentication off, so + // it is the already-disabled state that is the one-way one. + const sshAuth = guarded(config.disableSshAuth, (p) => p.disableSshAuth, true); const jwtTtlId = useId(); const [jwtTtlInput, setJwtTtlInput] = useState(String(config.sshJwtCacheTtl)); @@ -46,9 +85,11 @@ export function SettingsSSH() { setField("serverSshAllowed", v)} + disabled={sshServer.disabled} label={t("settings.ssh.server.label")} helpText={t("settings.ssh.server.help")} /> + {sshServer.hint} setField("enableSshRoot", v)} + disabled={sshRoot.disabled} label={t("settings.ssh.root.label")} helpText={t("settings.ssh.root.help")} /> + {sshRoot.hint} setField("enableSshSftp", v)} @@ -88,9 +131,11 @@ export function SettingsSSH() { setField("disableSshAuth", !v)} + disabled={sshAuth.disabled} label={t("settings.ssh.jwt.label")} helpText={t("settings.ssh.jwt.help")} /> + {sshAuth.hint}
); } + +// PrivilegeHint explains what an unprivileged user can and cannot do with a +// guarded control, and offers the command that does it with the privileges the +// daemon requires. oneWay covers the control being in the guarded state already: +// switching it back is the part that needs privileges. +function PrivilegeHint({ + actor, + command, + oneWay, + inverted, +}: { + actor: string; + command: string; + oneWay: boolean; + inverted: boolean; +}): ReactNode { + const { t } = useTranslation(); + if (!command) return null; + return ( +
+ + {!oneWay + ? t("settings.ssh.privilege.hint", { actor }) + : inverted + ? t("settings.ssh.privilege.oneWayInverted", { actor }) + : t("settings.ssh.privilege.oneWay", { actor })} + + + + {command} + + +
+ ); +} diff --git a/client/ui/grpc.go b/client/ui/grpc.go index c8e3aed76..5450e136d 100644 --- a/client/ui/grpc.go +++ b/client/ui/grpc.go @@ -5,14 +5,13 @@ package main import ( "fmt" "runtime" - "strings" "sync" "time" "google.golang.org/grpc" "google.golang.org/grpc/backoff" - "google.golang.org/grpc/credentials/insecure" + "github.com/netbirdio/netbird/client/internal/daemonaddr" "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/ui/desktop" ) @@ -36,9 +35,10 @@ func (c *Conn) Client() (proto.DaemonServiceClient, error) { return c.client, nil } - cc, err := grpc.NewClient( - strings.TrimPrefix(c.addr, "tcp://"), - grpc.WithTransportCredentials(insecure.NewCredentials()), + // Lazy on purpose: grpc.NewClient does not connect here, so a daemon that + // is down surfaces as a per-RPC Unavailable instead of blocking the UI. + target, opts := daemonaddr.DialTarget(daemonaddr.ResolveDaemonAddr(c.addr)) + opts = append(opts, grpc.WithUserAgent(desktop.GetUIUserAgent()), // Cap reconnect backoff at 5s; gRPC's default 120s MaxDelay would // leave the UI waiting 30-60s to notice a freshly-started daemon. @@ -51,6 +51,8 @@ func (c *Conn) Client() (proto.DaemonServiceClient, error) { }, }), ) + + cc, err := grpc.NewClient(target, opts...) if err != nil { return nil, fmt.Errorf("dial daemon: %w", err) } @@ -58,10 +60,12 @@ func (c *Conn) Client() (proto.DaemonServiceClient, error) { return c.client, nil } -// DaemonAddr returns the default daemon gRPC address: a Unix socket on Linux/macOS, TCP loopback on Windows. +// DaemonAddr returns the default daemon gRPC address: a Unix socket on +// Linux/macOS, a named pipe on Windows. The pipe carries the caller's token, +// which loopback TCP does not, so the daemon can tell who is calling. func DaemonAddr() string { if runtime.GOOS == "windows" { - return "tcp://127.0.0.1:41731" + return daemonaddr.WindowsPipeAddr } return "unix:///var/run/netbird.sock" } diff --git a/client/ui/i18n/locales/en/common.json b/client/ui/i18n/locales/en/common.json index 24bbc67ce..b668146e8 100644 --- a/client/ui/i18n/locales/en/common.json +++ b/client/ui/i18n/locales/en/common.json @@ -1774,5 +1774,17 @@ "error.unknown": { "message": "Operation failed.", "description": "Generic fallback error message used when no specific error applies." + }, + "settings.ssh.privilege.hint": { + "message": "Requires {actor}. Run this instead:", + "description": "Help text under an SSH setting the user cannot change: it needs elevated privileges. {actor} is 'root' on Linux/macOS or 'administrator privileges' on Windows. Followed by a copyable command." + }, + "settings.ssh.privilege.oneWay": { + "message": "You can switch this off, but switching it back on needs {actor}:", + "description": "Warning under an SSH setting an unprivileged user may disable but not re-enable. {actor} is 'root' on Linux/macOS or 'administrator privileges' on Windows. Followed by a copyable command." + }, + "settings.ssh.privilege.oneWayInverted": { + "message": "You can switch this on, but switching it back off needs {actor}:", + "description": "Warning under the SSH authentication setting, which an unprivileged user may re-enable but not disable again. {actor} is 'root' on Linux/macOS or 'administrator privileges' on Windows. Followed by a copyable command." } } diff --git a/client/ui/main.go b/client/ui/main.go index 9a3e17743..782e76d1c 100644 --- a/client/ui/main.go +++ b/client/ui/main.go @@ -96,7 +96,6 @@ func main() { } }) - settings := services.NewSettings(conn) profiles := services.NewProfiles(conn) // updater.Holder owns the typed update State; DaemonFeed feeds it and the // Update service is a thin Wails-bound facade over it plus the install RPCs. @@ -117,6 +116,7 @@ func main() { bundle, prefStore, localizer := buildI18n(app) // After bundle + prefStore: both are used to localise daemon errors. + settings := services.NewSettings(conn, bundle, prefStore, daemonAddr) connection := services.NewConnection(conn, bundle, prefStore) profileSwitcher := services.NewProfileSwitcher(profiles, connection, daemonFeed) // authsession.Session owns the full extend + dismiss surface the tray diff --git a/client/ui/services/errors.go b/client/ui/services/errors.go index f1679e764..0c6f2f20f 100644 --- a/client/ui/services/errors.go +++ b/client/ui/services/errors.go @@ -6,13 +6,30 @@ import ( "encoding/json" "strings" + "google.golang.org/genproto/googleapis/rpc/errdetails" gcodes "google.golang.org/grpc/codes" gstatus "google.golang.org/grpc/status" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/ui/i18n" "github.com/netbirdio/netbird/client/ui/preferences" ) +// privilegeErrorInfo returns the daemon's privilege-refusal detail, if the error +// carries one. +func privilegeErrorInfo(err error) (*errdetails.ErrorInfo, bool) { + for _, detail := range gstatus.Convert(err).Details() { + info, ok := detail.(*errdetails.ErrorInfo) + if !ok { + continue + } + if info.GetReason() == ipcauth.ErrorReasonPrivilegeRequired && info.GetDomain() == ipcauth.ErrorDomain { + return info, true + } + } + return nil, false +} + // ErrorTranslator localises daemon errors; runtime impl is *i18n.Bundle. type ErrorTranslator interface { Translate(lang i18n.LanguageCode, key string, args ...string) string @@ -30,6 +47,10 @@ type ClientError struct { Code string `json:"code"` Short string `json:"short"` Long string `json:"long"` + // Command is a command the user can run to complete the operation + // themselves, set when the daemon refused it for want of privileges. The + // frontend offers it for copying. + Command string `json:"command,omitempty"` } // Error returns the short message for plain Go callers. @@ -72,6 +93,24 @@ func (c errorClassifier) classify(err error) *ClientError { msg = st.Message() grpcCode = st.Code() } + + // A refusal for want of privileges carries its own summary and the command + // that performs the operation, both written for the user. Surface them + // verbatim: no substring guessing, and no localisation of a message the + // daemon composed. + if info, ok := privilegeErrorInfo(err); ok { + summary := info.GetMetadata()[ipcauth.ErrorMetaSummary] + if summary == "" { + summary = msg + } + return &ClientError{ + Code: "privilege_required", + Short: summary, + Long: summary, + Command: info.GetMetadata()[ipcauth.ErrorMetaCommand], + } + } + lower := strings.ToLower(msg) code := "unknown" diff --git a/client/ui/services/settings.go b/client/ui/services/settings.go index 3b6f6f81b..74e6f913c 100644 --- a/client/ui/services/settings.go +++ b/client/ui/services/settings.go @@ -7,6 +7,10 @@ import ( "fmt" "reflect" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/daemonaddr" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" ) @@ -39,6 +43,19 @@ type Restrictions struct { Features Features `json:"features"` } +// Privilege tells the frontend whether this process may perform the changes the +// daemon restricts to root/administrator, and carries the command for each so a +// disabled control can show the way to do it. +type Privilege struct { + Privileged bool `json:"privileged"` + // Actor names what the operation requires ("root", "administrator privileges"). + Actor string `json:"actor"` + // Commands equivalent to the settings the daemon guards, ready to copy. + AllowSSHServer string `json:"allowSshServer"` + EnableSSHRoot string `json:"enableSshRoot"` + DisableSSHAuth string `json:"disableSshAuth"` +} + type ConfigParams struct { ProfileName string `json:"profileName"` Username string `json:"username"` @@ -106,11 +123,19 @@ type SetConfigParams struct { } type Settings struct { - conn DaemonConn + conn DaemonConn + classifier errorClassifier + // daemonAddr is where the daemon listens, used to tell whether it runs as + // this user and would therefore authorize us: see Privilege. + daemonAddr string } -func NewSettings(conn DaemonConn) *Settings { - return &Settings{conn: conn} +func NewSettings(conn DaemonConn, translator ErrorTranslator, prefs LanguagePreference, daemonAddr string) *Settings { + return &Settings{ + conn: conn, + classifier: errorClassifier{translator: translator, prefs: prefs}, + daemonAddr: daemonAddr, + } } func (s *Settings) GetConfig(ctx context.Context, p ConfigParams) (Config, error) { @@ -189,8 +214,47 @@ func (s *Settings) SetConfig(ctx context.Context, p SetConfigParams) error { DisableSSHAuth: p.DisableSSHAuth, SshJWTCacheTTL: p.SSHJWTCacheTTL, } - _, err = cli.SetConfig(ctx, req) - return err + if _, err := cli.SetConfig(ctx, req); err != nil { + // Classified so the frontend gets the daemon's guidance instead of the + // gRPC envelope, which is what a refused privileged change looks like. + return s.classifier.classify(err) + } + return nil +} + +// Privilege reports whether this UI process could carry out the changes the +// daemon restricts to root/administrator, and the command that performs the one +// users hit in the SSH settings. It applies the daemon's own rule to what it can +// see locally, so the frontend can present those controls as unavailable up front +// instead of letting a save fail. No daemon round-trip, so it also works while the +// daemon is down. +// +// Being root or an elevated administrator is one way. The other is running as the +// daemon's own user while the daemon is unprivileged, which the daemon accepts +// because such a caller can already rewrite the config it reads; that is the +// rootless-container and Windows netstack-mode case, and it is read from the +// ownership of the socket or pipe the daemon created. +func (s *Settings) Privilege() Privilege { + id, err := ipcauth.CurrentProcessIdentity() + if err != nil { + // Fail closed: report unprivileged, which only ever disables controls. + log.Warnf("cannot read this process's identity, treating it as unprivileged: %v", err) + return newPrivilege(false) + } + if id.IsPrivileged() { + return newPrivilege(true) + } + return newPrivilege(daemonaddr.DaemonRunsAsSelf(s.daemonAddr)) +} + +func newPrivilege(privileged bool) Privilege { + return Privilege{ + Privileged: privileged, + Actor: ipcauth.PrivilegedActor(), + AllowSSHServer: ipcauth.UpCommand("--allow-server-ssh"), + EnableSSHRoot: ipcauth.UpCommand("--enable-ssh-root"), + DisableSSHAuth: ipcauth.UpCommand("--disable-ssh-auth"), + } } func (s *Settings) GetRestrictions(ctx context.Context) (Restrictions, error) { diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index bac9790b4..5f7aaa7bd 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -393,15 +393,17 @@ func (s *WindowManager) CloseWelcome() { } } -// OpenError shows the custom error dialog; title/message are pre-localised and ride in the -// start URL. A second error replaces the open one via SetURL. Singleton, destroyed on close. -func (s *WindowManager) OpenError(title, message string) { +// OpenError shows the custom error dialog; title/message/command are pre-localised +// and ride in the start URL. command is optional and, when set, is offered for +// copying so the user can run the operation the daemon refused. A second error +// replaces the open one via SetURL. Singleton, destroyed on close. +func (s *WindowManager) OpenError(title, message, command string) { if ShuttingDown() { return } s.mu.Lock() defer s.mu.Unlock() - startURL := errorDialogURL(title, message) + startURL := errorDialogURL(title, message, command) if s.errorDialog == nil { s.errorDialog = s.app.Window.NewWithOptions( DialogWindowOptions("error", s.title("window.title.error"), startURL, s.linuxIcon), @@ -601,8 +603,8 @@ func (s *WindowManager) getScreenBasedOnCursorPosition() *application.Screen { return nil } -// errorDialogURL builds the error window's start URL with title/message as escaped query params. -func errorDialogURL(title, message string) string { +// errorDialogURL builds the error window's start URL with title/message/command as escaped query params. +func errorDialogURL(title, message, command string) string { q := url.Values{} if title != "" { q.Set("title", title) @@ -610,6 +612,9 @@ func errorDialogURL(title, message string) string { if message != "" { q.Set("message", message) } + if command != "" { + q.Set("command", command) + } startURL := "/#/dialog/error" if enc := q.Encode(); enc != "" { startURL += "?" + enc diff --git a/go.mod b/go.mod index ca798decc..26a194958 100644 --- a/go.mod +++ b/go.mod @@ -30,6 +30,7 @@ require ( require ( github.com/DeRuina/timberjack v1.4.2 + github.com/Microsoft/go-winio v0.6.2 github.com/awnumar/memguard v0.23.0 github.com/aws/aws-sdk-go-v2 v1.38.3 github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.1 @@ -133,6 +134,7 @@ require ( golang.org/x/term v0.45.0 golang.org/x/time v0.15.0 google.golang.org/api v0.276.0 + google.golang.org/genproto/googleapis/rpc v0.0.0-20260401024825-9d38bb4040a9 gopkg.in/yaml.v3 v3.0.1 gorm.io/driver/mysql v1.5.7 gorm.io/driver/postgres v1.5.7 @@ -156,7 +158,6 @@ require ( github.com/Masterminds/goutils v1.1.1 // indirect github.com/Masterminds/semver/v3 v3.4.0 // indirect github.com/Masterminds/sprig/v3 v3.3.0 // indirect - github.com/Microsoft/go-winio v0.6.2 // indirect github.com/adrg/xdg v0.5.3 // indirect github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be // indirect github.com/apapsch/go-jsonmerge/v2 v2.0.0 // indirect @@ -317,7 +318,6 @@ require ( golang.org/x/tools v0.47.0 // indirect golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect google.golang.org/genproto/googleapis/api v0.0.0-20260319201613-d00831a3d3e7 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20260401024825-9d38bb4040a9 // indirect gopkg.in/square/go-jose.v2 v2.6.0 // indirect gopkg.in/tomb.v1 v1.0.0-20141024135613-dd632973f1e7 // indirect gopkg.in/yaml.v2 v2.4.0 // indirect From 1bf54ddd8f82476ed89fe6c9211f296e2ef284ff Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Thu, 30 Jul 2026 10:27:52 +0200 Subject: [PATCH 23/32] [client] Support Andorid session expiry handling (#6945) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Adds the session surface the Android client was missing: read the status label and session deadline, receive the engine's expiry warnings, extend the session via SSO without dropping the tunnel (cancellable), and dismiss a warning. Status() latches NeedsLogin so an engine restart doesn't erase it. ## Issue ticket number and link ## Stack ### Checklist - [ ] Is it a bug fix - [ ] Is a typo/documentation fix - [x] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ --- View with [code]smith Autofix with [code]smith Need help on this PR? Tag @codesmith-bot with what you need. Autofix is disabled. ## Summary by CodeRabbit * **New Features** * Added Android session status and session expiration details. * Added listener support for wake-state changes and session-expiry warnings. * Added interactive authentication session extension with async login-success handling, error reporting, warning dismissal, and cancellation of an in-progress extension. * **Bug Fixes** * Improved “login required” state handling after successful sign-in to prevent stale login-needed prompts. --- client/android/client.go | 25 ++- client/android/session.go | 309 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 333 insertions(+), 1 deletion(-) create mode 100644 client/android/session.go diff --git a/client/android/client.go b/client/android/client.go index b3a845818..1a8dd7d09 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -76,6 +76,24 @@ type Client struct { connectClient *internal.ConnectClient config *profilemanager.Config cacheDir string + + stateChangeMu sync.Mutex + stateChangeSubID string + eventSub *peer.EventSubscription + // Closed to stop the watch goroutines from delivering buffered items to a + // listener that has been removed or replaced. See stopStateChangeWatchLocked. + stateChangeDone chan struct{} + + // Latched "the server wants an interactive login": survives the engine + // restarts that replace the run loop's context state. See Client.Status. + // Guarded by loginRequiredMu together with loginCleared, which counts + // clears so a stale observation cannot re-latch over one. + loginRequiredMu sync.Mutex + loginRequired bool + loginCleared uint64 + + extendMu sync.Mutex + extendCancel context.CancelFunc } func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cc *internal.ConnectClient) { @@ -149,11 +167,16 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid if err != nil { return err } - // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) c.setState(cfg, cacheDir, connectClient) + // This path runs the interactive SSO flow, so reaching here means the peer + // is authenticated again — release the latch Status() reports from. Clear + // only once the fresh connect client is installed: until then Status() + // still reads the previous run's context state, which holds the NeedsLogin + // that prompted this login, and would re-latch what was just cleared. + c.clearLoginRequired() return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } diff --git a/client/android/session.go b/client/android/session.go new file mode 100644 index 000000000..961d52528 --- /dev/null +++ b/client/android/session.go @@ -0,0 +1,309 @@ +//go:build android + +package android + +import ( + "context" + "fmt" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal" + "github.com/netbirdio/netbird/client/internal/auth" + "github.com/netbirdio/netbird/client/internal/auth/sessionwatch" + "github.com/netbirdio/netbird/client/internal/peer" + cProto "github.com/netbirdio/netbird/client/proto" +) + +// StateChangeListener receives client state notifications. +// +// OnStateChanged is a payload-free wake-up whenever the state snapshot +// changed: connection state, the run-loop status label (e.g. NeedsLogin) or +// the session deadline. It mirrors the daemon's SubscribeStatus stream +// trigger — on each signal the consumer pulls the fresh values via +// Status() / SessionExpiresAtUnix(). +// +// OnSessionExpiring forwards the engine's session-expiry warnings, fired at +// sessionwatch.WarningLead before the deadline and again at FinalWarningLead +// (finalWarning true). The second one is suppressed when the user dismissed +// the first via DismissSessionWarning. The daemon turns the same events into +// its tray notification. +type StateChangeListener interface { + OnStateChanged() + OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool) +} + +// Status returns the connect run-loop's status label — the same value the +// desktop daemon serves in StatusResponse.Status. "NeedsLogin" means the +// management server rejected the peer and an interactive login is required. +// +// The label is latched: the run loop keeps its status in a per-run context +// state, which a restart replaces with a fresh Idle one, so an engine restart +// (network change, always-on) would otherwise erase the fact that the peer +// still needs to log in. Only a successful interactive login or extend clears +// it — see clearLoginRequired. +func (c *Client) Status() string { + latched, generation := c.loginRequiredState() + if latched { + return string(internal.StatusNeedsLogin) + } + cc := c.getConnectClient() + if cc == nil { + return string(internal.StatusIdle) + } + status := cc.Status() + if status == internal.StatusNeedsLogin { + c.latchLoginRequired(generation) + } + return string(status) +} + +func (c *Client) loginRequiredState() (bool, uint64) { + c.loginRequiredMu.Lock() + defer c.loginRequiredMu.Unlock() + return c.loginRequired, c.loginCleared +} + +// latchLoginRequired records a NeedsLogin observation, unless a clear landed +// while the caller was reading the run loop's status: cc.Status() is read +// outside the lock, so a login or extend completing in that window would +// otherwise be undone by this stale observation, stranding the UI on +// "login required" over a healthy session. +func (c *Client) latchLoginRequired(observedGeneration uint64) { + c.loginRequiredMu.Lock() + defer c.loginRequiredMu.Unlock() + if c.loginCleared != observedGeneration { + return + } + c.loginRequired = true +} + +// clearLoginRequired releases the latch after a successful interactive login +// or session extend, and invalidates any observation already in flight. +func (c *Client) clearLoginRequired() { + c.loginRequiredMu.Lock() + defer c.loginRequiredMu.Unlock() + c.loginRequired = false + c.loginCleared++ +} + +// SessionExpiresAtUnix returns the SSO session deadline as unix seconds, or 0 +// when no deadline is known (not SSO-registered, expiry disabled, or the +// engine has not received one yet). A past value means the session expired. +// Mirror of StatusResponse.sessionExpiresAt on the desktop daemon. +func (c *Client) SessionExpiresAtUnix() int64 { + deadline := c.recorder.GetSessionExpiresAt() + if deadline.IsZero() { + return 0 + } + return deadline.Unix() +} + +// SetStateChangeListener registers the state notification listener. +// Replaces any previously registered listener; remove it with +// RemoveStateChangeListener. +func (c *Client) SetStateChangeListener(listener StateChangeListener) { + c.stateChangeMu.Lock() + defer c.stateChangeMu.Unlock() + c.stopStateChangeWatchLocked() + if listener == nil { + return + } + + // Both subscriptions are buffered (one pending tick, ten pending events), + // so unsubscribing is not enough to stop callbacks: the loops would drain + // what is already queued and deliver it to a listener the caller has + // already removed or replaced. Gate every callback on this registration's + // own signal, which is closed before unsubscribing. + done := make(chan struct{}) + c.stateChangeDone = done + + id, ch := c.recorder.SubscribeToStateChanges() + c.stateChangeSubID = id + // The channel is closed by UnsubscribeFromStateChanges, which ends the + // goroutine. Ticks are coalesced (buffer of one), so a burst of changes + // wakes the listener once. + go func() { + for range ch { + select { + case <-done: + return + default: + } + listener.OnStateChanged() + } + }() + + c.eventSub = c.recorder.SubscribeToEvents() + go watchSessionWarnings(c.eventSub, listener, done) +} + +// RemoveStateChangeListener unregisters the state notification listener. +func (c *Client) RemoveStateChangeListener() { + c.stateChangeMu.Lock() + defer c.stateChangeMu.Unlock() + c.stopStateChangeWatchLocked() +} + +// DismissSessionWarning records the user's "Dismiss" on the first expiry +// warning and suppresses the final one for the current deadline. A refreshed +// deadline re-arms both. No-op while the engine is not running. +func (c *Client) DismissSessionWarning() { + cc := c.getConnectClient() + if cc == nil { + return + } + engine := cc.Engine() + if engine == nil { + return + } + engine.DismissSessionWarning() +} + +// ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and +// asks the management server to extend the session deadline. The tunnel is +// untouched: no resync, no reconnect. Async; the result arrives on the +// listener. Mirror of the daemon's RequestExtendAuthSession / +// WaitExtendAuthSession RPC pair, with URLOpener playing the "UI opens the +// browser" role. +// +// Only one flow may be in flight: the PKCE step binds a fixed loopback port, +// so a second concurrent flow would fail on that bind. Call +// CancelExtendAuthSession when the user abandons the browser. +func (c *Client) ExtendAuthSession(urlOpener URLOpener, isAndroidTV bool, resultListener ErrListener) { + ctx, err := c.beginExtend() + if err != nil { + resultListener.OnError(err) + return + } + + go func() { + defer c.endExtend() + if err := c.extendAuthSession(ctx, urlOpener, isAndroidTV); err != nil { + resultListener.OnError(err) + return + } + resultListener.OnSuccess() + }() +} + +// CancelExtendAuthSession aborts an in-flight ExtendAuthSession. The tunnel is +// left alone — unlike the login flow, which cancels the whole client context +// by stopping the engine. Without this the abandoned PKCE wait keeps its +// loopback port for the full flow timeout and blocks every later attempt. +// No-op when no flow is running. +func (c *Client) CancelExtendAuthSession() { + c.extendMu.Lock() + defer c.extendMu.Unlock() + if c.extendCancel != nil { + c.extendCancel() + } +} + +func (c *Client) stopStateChangeWatchLocked() { + // Signal first, unsubscribe second: closing the channels only stops new + // items, and the loops would still hand whatever is buffered to a listener + // that is no longer registered. + if c.stateChangeDone != nil { + close(c.stateChangeDone) + c.stateChangeDone = nil + } + if c.stateChangeSubID != "" { + c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID) + c.stateChangeSubID = "" + } + if c.eventSub != nil { + // Closes the channel, which ends watchSessionWarnings. + c.recorder.UnsubscribeFromEvents(c.eventSub) + c.eventSub = nil + } +} + +// watchSessionWarnings forwards the engine's session-expiry warnings to the +// listener. The event stream also carries unrelated traffic — network-map +// updates on every sync, DNS and route errors — so everything but an +// AUTHENTICATION event carrying the session-warning marker is dropped. Exits +// when the subscription is closed by UnsubscribeFromEvents, or earlier when +// done is closed — the stream buffers up to ten events, and a deregistered +// listener must not receive the ones already queued. +func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) { + for ev := range sub.Events() { + select { + case <-done: + return + default: + } + if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION { + continue + } + meta := ev.GetMetadata() + if meta[sessionwatch.MetaSessionWarning] != "true" { + // Other AUTHENTICATION events exist (e.g. a deadline rejected as + // out of range); they carry no warning marker. + continue + } + deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt]) + if err != nil { + log.Warnf("session warning event with unparsable deadline: %v", err) + continue + } + lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes]) + if err != nil { + // Informational only — the deadline above is what drives the UI. + lead = 0 + } + listener.OnSessionExpiring(deadline.Unix(), int64(lead), + meta[sessionwatch.MetaSessionFinal] == "true") + } +} + +func (c *Client) beginExtend() (context.Context, error) { + c.extendMu.Lock() + defer c.extendMu.Unlock() + if c.extendCancel != nil { + return nil, fmt.Errorf("session extend already in progress") + } + ctx, cancel := context.WithCancel(context.Background()) + c.extendCancel = cancel + return ctx, nil +} + +func (c *Client) endExtend() { + c.extendMu.Lock() + defer c.extendMu.Unlock() + if c.extendCancel != nil { + c.extendCancel() + c.extendCancel = nil + } +} + +func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isAndroidTV bool) error { + cfg, _, cc := c.stateSnapshot() + if cfg == nil || cc == nil { + return fmt.Errorf("engine is not running") + } + engine := cc.Engine() + if engine == nil { + return fmt.Errorf("engine is not initialized") + } + + authClient, err := auth.NewAuth(ctx, cfg.PrivateKey, cfg.ManagementURL, cfg) + if err != nil { + return fmt.Errorf("failed to create auth client: %v", err) + } + defer authClient.Close() + + a := &Auth{ctx: ctx, config: cfg} + tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV) + if err != nil { + return fmt.Errorf("interactive sso login failed: %v", err) + } + + if _, err := engine.ExtendAuthSession(ctx, tokenInfo.GetTokenToUse()); err != nil { + return err + } + c.clearLoginRequired() + + go urlOpener.OnLoginSuccess() + return nil +} From cff49237b620d5eff09092a632220e43897c975d Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Thu, 30 Jul 2026 20:11:27 +0900 Subject: [PATCH 24/32] [client] Stop and remove the daemon on netbird-ui cask uninstall (#6977) --- client/ui/netbird-ui.rb.tmpl | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/client/ui/netbird-ui.rb.tmpl b/client/ui/netbird-ui.rb.tmpl index 06971909d..1c77e6717 100644 --- a/client/ui/netbird-ui.rb.tmpl +++ b/client/ui/netbird-ui.rb.tmpl @@ -29,8 +29,13 @@ cask "{{ $projectName }}" do end uninstall_preflight do - system_command "#{appdir}/Netbird UI.app/uninstaller.sh", - sudo: false + system_command "/bin/sh", + args: ["-c", <<~CMD], + launchctl bootout system/netbird 2>/dev/null || \ + launchctl unload /Library/LaunchDaemons/netbird.plist 2>/dev/null || true + rm -f /Library/LaunchDaemons/netbird.plist + CMD + sudo: true end name "Netbird UI" From c1f0006012cf1153723a3bd094caf9881e2d09e8 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Thu, 30 Jul 2026 22:34:08 +0900 Subject: [PATCH 25/32] [misc] Update SECURITY.md (#6981) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes ## Issue ticket number and link ## Stack ### Checklist - [ ] Is it a bug fix - [x] Is a typo/documentation fix - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) - [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why) ### Docs PR URL (required if "docs added" is checked) Paste the PR link from https://github.com/netbirdio/docs here: https://github.com/netbirdio/docs/pull/__ ## Summary by CodeRabbit * **Documentation** * Expanded the security vulnerability reporting policy with private reporting options and guidance for hosted infrastructure issues. * Added recommendations for report contents, acknowledgements, severity assessment, remediation, advisories, and reporter credit. * Clarified supported versions, advisory distribution, bug bounty status, and where to report non-security issues. --- SECURITY.md | 72 +++++++++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 65 insertions(+), 7 deletions(-) diff --git a/SECURITY.md b/SECURITY.md index 745c66e61..bdf88d670 100644 --- a/SECURITY.md +++ b/SECURITY.md @@ -1,12 +1,70 @@ # Security Policy -NetBird's goal is to provide a secure network. If you find a vulnerability or bug, please report it by opening an issue [here](https://github.com/netbirdio/netbird/issues/new?assignees=&labels=&template=bug-issue-report.md&title=) or by contacting us by email. - -There has yet to be an official bug bounty program for the NetBird project. - -## Supported Versions -- We currently support only the latest version +NetBird's goal is to provide a secure network. The client runs as a privileged service on every machine it is installed on, +so we take reports about it seriously and we publish what we fix. ## Reporting a Vulnerability -Please report security issues to `security@netbird.io` +**Please do not open a public issue for a security vulnerability.** Public issues are visible to everyone, including before +a fix is available. + +Report security issues one of these two ways: + +- **GitHub private vulnerability reporting** — [open a private report](https://github.com/netbirdio/netbird/security/advisories/new) + on this repository. This is the preferred route: it keeps the discussion, the draft advisory, and the credit in one place. +- **Email** — `security@netbird.io`. + +If the finding affects NetBird Cloud or our hosted infrastructure rather than the open-source code, email us rather than +filing a repository report. + +### What to include + +A report is easier to act on when it contains: + +- The affected component (client, management, signal, relay, dashboard) and the version or commit you tested +- The platform and configuration, where relevant — operating system, self-hosted or NetBird Cloud, container or host install +- What an attacker needs before they can exploit it: network position, an account, local access, a specific privilege level +- Steps to reproduce, and a proof of concept if you have one +- The impact you believe it has + +Partial reports are still welcome. If you are unsure whether something is a security issue, send it to `security@netbird.io` +and let us make that call. + +## What to expect from us + +- **We acknowledge your report** and tell you whether we can reproduce it. +- **We work with you on severity and scope.** If we assess it differently than you do, we will explain why rather than + silently downgrade it. +- **We fix and release**, then publish a [GitHub Security Advisory](https://github.com/netbirdio/netbird/security/advisories) + naming the affected version range and the patched version. +- **We credit reporters who want to be credited.** Tell us the name or handle you would like used, or that you would rather + stay anonymous. +- **We keep you in the loop** until the advisory is published. + +We ask that you give us a reasonable opportunity to ship a fix before disclosing the issue publicly, and that you avoid +accessing, modifying, or exfiltrating data belonging to other people while testing. Testing against your own installation +or your own account is always fine. + +## Supported Versions + +We support the latest release. Security fixes ship in the next version rather than as backports to older releases, so +upgrading to the current release is how you get them. + +Release notifications are available by watching [releases](https://github.com/netbirdio/netbird/releases). + +## Published advisories + +Every vulnerability we fix is published as a GitHub Security Advisory on the +[advisories page](https://github.com/netbirdio/netbird/security/advisories), including the affected version range, the +patched version, and the reporter's credit. Advisories for the Go module are also distributed through the Go vulnerability +database, so `govulncheck` will report them against your dependencies. + +## Bug bounty + +There is no official bug bounty program for the NetBird project. We credit reporters in advisories, and we are grateful for +the work, but we cannot currently offer payment for reports. + +## Non-security bugs + +For bugs that are not security issues, please use the +[issue tracker](https://github.com/netbirdio/netbird/discussions/new/choose). From 234abd7a083d6501413692cc8bd2631c3375bf72 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Fri, 31 Jul 2026 12:34:58 +0900 Subject: [PATCH 26/32] [misc] group x package updates and run weekly (#7000) --- .github/dependabot.yml | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/.github/dependabot.yml b/.github/dependabot.yml index b78b1417a..647e04936 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -3,8 +3,8 @@ updates: - package-ecosystem: "github-actions" directory: "/" schedule: - interval: "daily" - open-pull-requests-limit: 15 + interval: "weekly" + open-pull-requests-limit: 3 groups: actions: patterns: @@ -22,9 +22,12 @@ updates: directories: - "/" schedule: - interval: "daily" + interval: "weekly" open-pull-requests-limit: 15 groups: + golang-x-packages: + patterns: + - "golang.org/x/*" aws-sdk: patterns: - "github.com/aws/aws-sdk-go-v2/*" From aed60a24323c725ed1e88fd7fcae1cdfb211a9d5 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Fri, 31 Jul 2026 22:53:58 +0900 Subject: [PATCH 27/32] [client] Fix daemon lock order inversion between SetConfig and login (#6978) --- client/server/lock_order_test.go | 51 ++++++++++++++++++++++++++++++++ client/server/server.go | 16 ++++++---- 2 files changed, 61 insertions(+), 6 deletions(-) create mode 100644 client/server/lock_order_test.go diff --git a/client/server/lock_order_test.go b/client/server/lock_order_test.go new file mode 100644 index 000000000..457e3db34 --- /dev/null +++ b/client/server/lock_order_test.go @@ -0,0 +1,51 @@ +package server + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/proto" +) + +// The daemon takes guardedConfigMu before s.mutex. authorizeAndPrepareLogin +// takes s.mutex while holding guardedConfigMu, so a SetConfig that grabbed +// s.mutex first and then waited for guardedConfigMu would deadlock the daemon +// against a concurrent login: two unprivileged IPC calls are enough. +// +// The held guardedConfigMu below stands in for that login. While SetConfig waits +// for it, s.mutex must stay free, otherwise the login waiting for s.mutex could +// never release guardedConfigMu. +func TestSetConfig_TakesGuardedConfigMuBeforeServerMutex(t *testing.T) { + s, ctx, profName, username, _ := setupServerWithProfile(t) + + s.guardedConfigMu.Lock() + + done := make(chan error, 1) + go func() { + _, err := s.SetConfig(ctx, &proto.SetConfigRequest{ + ProfileName: profName, + Username: username, + }) + done <- err + }() + + require.Never(t, func() bool { + if !s.mutex.TryLock() { + return true + } + s.mutex.Unlock() + return false + }, 500*time.Millisecond, 10*time.Millisecond, + "SetConfig held s.mutex while waiting for guardedConfigMu, which deadlocks against a concurrent login") + + s.guardedConfigMu.Unlock() + + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("SetConfig did not finish after guardedConfigMu was released") + } +} diff --git a/client/server/server.go b/client/server/server.go index db5909272..6e22a76a9 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -393,6 +393,16 @@ func (s *Server) loginAttempt(ctx context.Context, setupKey, jwtToken string) (i // Login uses setup key to prepare configuration for the daemon. func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigRequest) (*proto.SetConfigResponse, error) { + // Privilege gate: refuse the parts of the request that would let a local + // user turn the root daemon into a root shell. Held across the write so the + // config cannot gain the SSH server between the decision and the update. + // + // Taken before s.mutex: authorizeAndPrepareLogin takes s.mutex while holding + // guardedConfigMu, so acquiring the two in the other order here would let a + // concurrent login deadlock the daemon. + s.guardedConfigMu.Lock() + defer s.guardedConfigMu.Unlock() + s.mutex.Lock() defer s.mutex.Unlock() @@ -417,12 +427,6 @@ func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigReques return nil, err } - // Privilege gate: refuse the parts of the request that would let a local - // user turn the root daemon into a root shell. Held across the write so the - // config cannot gain the SSH server between the decision and the update. - s.guardedConfigMu.Lock() - defer s.guardedConfigMu.Unlock() - stored, err := s.storedProfileConfig(msg.ProfileName, msg.Username) if err != nil { return nil, err From 7516aa6473cb1eba912ddd1f88e3563f5c671154 Mon Sep 17 00:00:00 2001 From: Ben Date: Fri, 31 Jul 2026 16:28:42 +0200 Subject: [PATCH 28/32] [client] Fix expression order in legacy nftables route rules (#7011) --- .../nftables/legacy_rule_linux_test.go | 60 +++++++++++++++++++ client/firewall/nftables/router_linux.go | 21 ++++--- 2 files changed, 72 insertions(+), 9 deletions(-) create mode 100644 client/firewall/nftables/legacy_rule_linux_test.go diff --git a/client/firewall/nftables/legacy_rule_linux_test.go b/client/firewall/nftables/legacy_rule_linux_test.go new file mode 100644 index 000000000..dc2f1c7a0 --- /dev/null +++ b/client/firewall/nftables/legacy_rule_linux_test.go @@ -0,0 +1,60 @@ +package nftables + +import ( + "testing" + + "github.com/google/nftables/expr" + "github.com/stretchr/testify/require" +) + +func TestBuildLegacyRouteRuleExpressions(t *testing.T) { + sourcePayload := &expr.Payload{} + sourceCmp := &expr.Cmp{} + destinationPayload := &expr.Payload{} + destinationCmp := &expr.Cmp{} + nilSourceDestination := &expr.Payload{} + nilDestinationSource := &expr.Cmp{} + + tests := []struct { + name string + source []expr.Any + destination []expr.Any + matches []expr.Any + }{ + { + name: "both non-empty", + source: []expr.Any{sourcePayload, sourceCmp}, + destination: []expr.Any{destinationPayload, destinationCmp}, + matches: []expr.Any{sourcePayload, sourceCmp, destinationPayload, destinationCmp}, + }, + { + name: "nil source", + destination: []expr.Any{nilSourceDestination}, + matches: []expr.Any{nilSourceDestination}, + }, + { + name: "nil destination", + source: []expr.Any{nilDestinationSource}, + matches: []expr.Any{nilDestinationSource}, + }, + { + name: "both nil", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := buildLegacyRouteRuleExpressions(tt.source, tt.destination) + + require.Len(t, got, len(tt.matches)+2) + for i, match := range tt.matches { + require.Same(t, match, got[i]) + } + + require.IsType(t, &expr.Counter{}, got[len(tt.matches)]) + verdict, ok := got[len(tt.matches)+1].(*expr.Verdict) + require.True(t, ok) + require.Equal(t, expr.VerdictAccept, verdict.Kind) + }) + } +} diff --git a/client/firewall/nftables/router_linux.go b/client/firewall/nftables/router_linux.go index 4214455a9..dfb94c514 100644 --- a/client/firewall/nftables/router_linux.go +++ b/client/firewall/nftables/router_linux.go @@ -953,6 +953,17 @@ func (r *router) addMSSClampingRules() error { return r.conn.Flush() } +func buildLegacyRouteRuleExpressions(sourceExp, destExp []expr.Any) []expr.Any { + exprs := make([]expr.Any, 0, len(sourceExp)+len(destExp)+2) + exprs = append(exprs, sourceExp...) + exprs = append(exprs, destExp...) + exprs = append(exprs, + &expr.Counter{}, + &expr.Verdict{Kind: expr.VerdictAccept}, + ) + return exprs +} + // addLegacyRouteRule adds a legacy routing rule for mgmt servers pre route acls func (r *router) addLegacyRouteRule(pair firewall.RouterPair) error { sourceExp, err := r.applyNetwork(pair.Source, nil, true) @@ -965,15 +976,7 @@ func (r *router) addLegacyRouteRule(pair firewall.RouterPair) error { return fmt.Errorf("apply destination: %w", err) } - exprs := []expr.Any{ - &expr.Counter{}, - &expr.Verdict{ - Kind: expr.VerdictAccept, - }, - } - - exprs = append(exprs, sourceExp...) - exprs = append(exprs, destExp...) + exprs := buildLegacyRouteRuleExpressions(sourceExp, destExp) ruleKey := firewall.GenKey(firewall.ForwardingFormat, pair) From aad2702a14f10121647ae70a2fb7093230e8def1 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Fri, 31 Jul 2026 16:56:41 +0200 Subject: [PATCH 29/32] [proxy] remove cluster tag from proxy metrics (#6985) --- .../reverseproxy/proxy/manager/controller.go | 6 ++--- .../reverseproxy/proxy/manager/metrics.go | 22 +++++-------------- 2 files changed, 9 insertions(+), 19 deletions(-) diff --git a/management/internals/modules/reverseproxy/proxy/manager/controller.go b/management/internals/modules/reverseproxy/proxy/manager/controller.go index e5b3e9886..0e5064f63 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/controller.go +++ b/management/internals/modules/reverseproxy/proxy/manager/controller.go @@ -36,7 +36,7 @@ func NewGRPCController(proxyGRPCServer *nbgrpc.ProxyServiceServer, meter metric. // SendServiceUpdateToCluster sends a service update to a specific proxy cluster. func (c *GRPCController) SendServiceUpdateToCluster(ctx context.Context, accountID string, update *proto.ProxyMapping, clusterAddr string) { c.proxyGRPCServer.SendServiceUpdateToCluster(ctx, update, clusterAddr) - c.metrics.IncrementServiceUpdateSendCount(clusterAddr) + c.metrics.IncrementServiceUpdateSendCount() } // GetOIDCValidationConfig returns the OIDC validation configuration from the gRPC server. @@ -53,7 +53,7 @@ func (c *GRPCController) RegisterProxyToCluster(ctx context.Context, clusterAddr proxySet.(*sync.Map).Store(proxyID, struct{}{}) log.WithContext(ctx).Debugf("Registered proxy %s to cluster %s", proxyID, clusterAddr) - c.metrics.IncrementProxyConnectionCount(clusterAddr) + c.metrics.IncrementProxyConnectionCount() return nil } @@ -67,7 +67,7 @@ func (c *GRPCController) UnregisterProxyFromCluster(ctx context.Context, cluster proxySet.(*sync.Map).Delete(proxyID) log.WithContext(ctx).Debugf("Unregistered proxy %s from cluster %s", proxyID, clusterAddr) - c.metrics.DecrementProxyConnectionCount(clusterAddr) + c.metrics.DecrementProxyConnectionCount() } return nil } diff --git a/management/internals/modules/reverseproxy/proxy/manager/metrics.go b/management/internals/modules/reverseproxy/proxy/manager/metrics.go index 2b402cead..f7fd385cb 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/metrics.go +++ b/management/internals/modules/reverseproxy/proxy/manager/metrics.go @@ -3,7 +3,6 @@ package manager import ( "context" - "go.opentelemetry.io/otel/attribute" "go.opentelemetry.io/otel/metric" ) @@ -48,25 +47,16 @@ func newMetrics(meter metric.Meter) (*metrics, error) { }, nil } -func (m *metrics) IncrementProxyConnectionCount(clusterAddr string) { - m.proxyConnectionCount.Add(context.Background(), 1, - metric.WithAttributes( - attribute.String("cluster", clusterAddr), - )) +func (m *metrics) IncrementProxyConnectionCount() { + m.proxyConnectionCount.Add(context.Background(), 1) } -func (m *metrics) DecrementProxyConnectionCount(clusterAddr string) { - m.proxyConnectionCount.Add(context.Background(), -1, - metric.WithAttributes( - attribute.String("cluster", clusterAddr), - )) +func (m *metrics) DecrementProxyConnectionCount() { + m.proxyConnectionCount.Add(context.Background(), -1) } -func (m *metrics) IncrementServiceUpdateSendCount(clusterAddr string) { - m.serviceUpdateSendCount.Add(context.Background(), 1, - metric.WithAttributes( - attribute.String("cluster", clusterAddr), - )) +func (m *metrics) IncrementServiceUpdateSendCount() { + m.serviceUpdateSendCount.Add(context.Background(), 1) } func (m *metrics) IncrementProxyHeartbeatCount() { From f51fadf8d4182c9d5ba4c99bdf392d7c03ba4ff2 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Sat, 1 Aug 2026 00:15:39 +0900 Subject: [PATCH 30/32] [misc] update contributing guide (#7009) ## Describe your changes We are updating the contributing guide to require an issue to be opened before PRs; this will allow discussing changes before code is shipped. Update the rest of the file because of outdated information and align the pull request template ## Issue ticket number and link --- .github/pull_request_template.md | 10 +- CONTRIBUTING.md | 287 ++++++++++++++++++++++++++++--- 2 files changed, 268 insertions(+), 29 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 8e68054bd..9b796f262 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -2,6 +2,12 @@ ## Issue ticket number and link + + ## Stack @@ -12,7 +18,9 @@ - [ ] Is a feature enhancement - [ ] It is a refactor - [ ] Created tests that fail without the change (if possible) -- [ ] This change does **not** modify the public API, gRPC protocols, functionality behavior, CLI / service flags, or introduce a new feature — **OR** I have discussed it with the NetBird team beforehand (link the issue / Slack thread in the description). See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first). +- [ ] I ran and tested this change locally — I did not rely on CI to find out whether it works +- [ ] This PR has a single purpose (not a fix + refactor + feature in one) +- [ ] This change is a trivial fix, **OR** it links an issue the NetBird team agreed on beforehand. Changes to the public API, gRPC protocols, functionality behavior, CLI / service flags, or new features always need that agreement first. See [CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second). > By submitting this pull request, you confirm that you have read and agree to the terms of the [Contributor License Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md). diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index d9c0b416e..3b8017788 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -1,6 +1,6 @@ # Contributing to NetBird -Thanks for your interest in contributing to NetBird. +Thanks for your interest in contributing to NetBird. There are many ways that you can contribute: - Reporting issues @@ -10,12 +10,69 @@ There are many ways that you can contribute: If you haven't already, join our slack workspace [here](https://docs.netbird.io/slack-url), we would love to discuss topics that need community contribution and enhancements to existing features. +## Ticket first, PR second + +**Open a ticket and wait for feedback before you open a pull request.** Every PR +that changes behavior must link to an issue the NetBird team has agreed on. A PR +that arrives without one may be closed and redirected to a discussion, no matter +how good the code is. + +Issues in this repository are maintainer-curated work items, so the flow starts +in [Discussions](https://github.com/netbirdio/netbird/discussions): + +1. **Open a discussion.** Use + [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) + for a bug, regression, or unexpected behavior, and + [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) + for a feature, enhancement, or integration idea. Setup and usage questions + belong in + [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support). + Never report a security vulnerability in public — follow the + [security policy](https://github.com/netbirdio/netbird/security/policy) + instead. +2. **Wait for feedback.** DevRel validates and reproduces the report, and a + maintainer confirms the direction. We may ask for more detail or propose a + different approach. Validated discussions become issues. +3. **Then write the code**, following the approach agreed in the issue, and open + the PR linking that issue. + +Trivial fixes — a typo, a broken link, a documentation correction, or a one-line +fix that already has an issue — can go straight to a PR. Everything else starts +with a ticket. When in doubt, ask in the discussion or on +[Slack](https://docs.netbird.io/slack-url); an hour of conversation up front +regularly saves a week of rework. + +### High-risk areas + +These always need the design discussed and agreed in the issue **before** you +write code: + +- **Public API** — REST / management API, OpenAPI schema, dashboard-facing contracts +- **gRPC protocols** — management, signal, relay, and client daemon protos +- **Functionality behavior** — anything existing deployments would experience differently after an upgrade +- **Peer connectivity** — ICE and NAT traversal, relay selection, WireGuard® and Rosenpass key handling +- **Client system integration** — routing, firewall, DNS, and interface management +- **Authentication and authorization** — IdP integration, tokens, permissions, cryptography +- **CLI / service flags**, configuration file format, and daemon IPC +- **Store and database schema** — models and migrations +- **New features** + +These surfaces are NetBird's contract with operators, self-hosters, and +downstream integrators, and changes to them have compatibility, security, and +release-planning implications. Agreeing on the direction early lets the PR +review focus on implementation rather than design. + +Typical bug fixes, internal refactors, documentation updates, and tests do not +need a design discussion, but should still be tied to an issue so the work is +visible and nobody duplicates it. + ## Contents - [Contributing to NetBird](#contributing-to-netbird) + - [Ticket first, PR second](#ticket-first-pr-second) + - [High-risk areas](#high-risk-areas) - [Contents](#contents) - [Code of conduct](#code-of-conduct) - - [Discuss changes with the NetBird team first](#discuss-changes-with-the-netbird-team-first) - [Directory structure](#directory-structure) - [Development setup](#development-setup) - [Requirements](#requirements) @@ -24,6 +81,7 @@ If you haven't already, join our slack workspace [here](https://docs.netbird.io/ - [Build and start](#build-and-start) - [Test suite](#test-suite) - [Checklist before submitting a PR](#checklist-before-submitting-a-pr) + - [When we close a PR](#when-we-close-a-pr) - [Other project repositories](#other-project-repositories) - [Contributor License Agreement](#contributor-license-agreement) @@ -34,42 +92,66 @@ Conduct which can be found in the file [CODE_OF_CONDUCT.md](CODE_OF_CONDUCT.md). By participating, you are expected to uphold this code. Please report unacceptable behavior to community@netbird.io. -## Discuss changes with the NetBird team first - -Changes to the **public API**, **gRPC protocols**, **functionality behavior**, **CLI / service flags**, or **new features** should be discussed with the NetBird team before you start the work. These surfaces are part of NetBird's contract with operators, self-hosters, and downstream integrators, and changes to them have compatibility, security, and release-planning implications that benefit from an early conversation. - -Open an issue or reach out on [Slack](https://docs.netbird.io/slack-url) to talk through what you have in mind. We'll help shape the change, flag any constraints we know about, and confirm the direction so the PR review can focus on implementation rather than design. - -Typical bug fixes, internal refactors, documentation updates, and tests do not need pre-discussion — open the PR directly. - ## Directory structure -The NetBird project monorepo is organized to maintain most of its individual dependencies code within their directories, except for a few auxiliary or shared packages. +The NetBird project monorepo keeps most of each component's code within its own +directory, except for a few auxiliary or shared packages. Protocol definitions +and the client-side service clients live under [/shared](/shared), because both +the agent and the services import them. -The most important directories are: +**Agent** -- [/.github](/.github) - Github actions workflow files and issue templates - [/client](/client) - NetBird agent code -- [/client/cmd](/client/cmd) - NetBird agent cli code +- [/client/cmd](/client/cmd) - NetBird agent CLI code - [/client/internal](/client/internal) - NetBird agent business logic code -- [/client/proto](/client/proto) - NetBird agent daemon GRPC proto files - [/client/server](/client/server) - NetBird agent daemon code for background execution -- [/client/ui](/client/ui) - NetBird agent UI code -- [/encryption](/encryption) - Contain main encryption code for agent communication -- [/iface](/iface) - Wireguard® interface code -- [/infrastructure_files](/infrastructure_files) - Getting started files containing docker and template scripts +- [/client/proto](/client/proto) - NetBird agent daemon gRPC proto files +- [/client/iface](/client/iface) - WireGuard® interface code +- [/client/firewall](/client/firewall) - Platform firewall backends (nftables, iptables, pf, WFP, userspace) +- [/client/ssh](/client/ssh) - Built-in SSH server and client +- [/client/ui](/client/ui) - NetBird agent UI code (Wails v3 + React) +- [/client/android](/client/android), [/client/ios](/client/ios) - Mobile platform bindings +- [/client/wasm](/client/wasm) - WebAssembly build of the agent +- [/client/mdm](/client/mdm) - MDM-delivered policy handling +- [/client/system](/client/system) - Host and system information collection + +**Control plane services** + - [/management](/management) - Management service code -- [/management/client](/management/client) - Management service client code which is imported by the agent code -- [/management/proto](/management/proto) - Management service GRPC proto files - [/management/server](/management/server) - Management service server code - [/management/server/http](/management/server/http) - Management service REST API code +- [/management/server/store](/management/server/store) - Persistence layer and migrations - [/management/server/idp](/management/server/idp) - Management service IDP management code -- [/release_files](/release_files) - Files that goes into release packages +- [/management/server/peer](/management/server/peer), [/management/server/groups](/management/server/groups), [/management/server/networks](/management/server/networks), [/management/server/posture](/management/server/posture), [/management/server/permissions](/management/server/permissions) - Core domain packages - [/signal](/signal) - Signal service code -- [/signal/client](/signal/client) - Signal service client code which is imported by the agent code - [/signal/peer](/signal/peer) - Signal service peer message logic -- [/signal/proto](/signal/proto) - Signal service GRPC proto files - [/signal/server](/signal/server) - Signal service server code +- [/relay](/relay) - Relay service code +- [/relay/protocol](/relay/protocol) - Relay wire protocol +- [/proxy](/proxy) - Identity-aware proxy used by Agent Network (LLM routing, ACME, access logs) +- [/agent-network](/agent-network) - Agent Network overview and documentation +- [/upload-server](/upload-server) - Debug bundle upload service + +**Shared code** + +- [/shared/management/proto](/shared/management/proto) - Management service gRPC proto files +- [/shared/management/client](/shared/management/client) - Management service client code which is imported by the agent code +- [/shared/management/http/api](/shared/management/http/api) - OpenAPI specification and generated REST API types +- [/shared/signal/proto](/shared/signal/proto) - Signal service gRPC proto files +- [/shared/signal/client](/shared/signal/client) - Signal service client code which is imported by the agent code +- [/shared/relay](/shared/relay) - Relay client and shared relay types +- [/shared/auth](/shared/auth), [/shared/sshauth](/shared/sshauth) - Shared authentication primitives +- [/encryption](/encryption) - Contain main encryption code for agent communication +- [/dns](/dns), [/route](/route), [/stun](/stun), [/sharedsock](/sharedsock), [/util](/util) - Shared networking and utility primitives +- [/flow](/flow) - Flow event protocol shared by the agent and Management + +**Build, test, and packaging** + +- [/.github](/.github) - Github actions workflow files, issue templates, and the pull request template +- [/e2e](/e2e) - End-to-end test suites and harness +- [/infrastructure_files](/infrastructure_files) - Getting started files containing docker and template scripts +- [/release_files](/release_files) - Files that goes into release packages +- [/tools](/tools) - Development and maintenance tooling ## Development setup @@ -334,23 +416,172 @@ The installer `netbird-installer.exe` will be created in root directory. ### Test suite -The tests can be started via: +The host-safe unit tests run as a normal user and leave host networking +untouched: ``` -cd netbird -go test -exec sudo ./... +make test-unit ``` + +Tests that need root and mutate host networking (firewall, routing, interface +management) carry the `privileged` build tag and run inside a +`--privileged --cap-add=NET_ADMIN` Docker container: + +``` +make test-privileged +``` + +Narrow a privileged run with environment variables: + +``` +PRIV_RUN=TestNftablesManager PRIV_PKGS=./client/firewall/nftables/... make test-privileged +``` + +Single packages can be run directly, adding `-race` when the change touches +shared state: + +``` +go test -race ./client/internal/dns/... +``` + > On Windows use a powershell with administrator privileges ## Checklist before submitting a PR -As a critical network service and open-source project, we must enforce a few things before submitting the pull-requests: + +As a critical network service and open-source project, we must enforce a few +things before submitting a pull request. The +[pull request template](/.github/pull_request_template.md) mirrors this list — +fill it in rather than deleting it. + +### Link the issue + +The PR description must link the agreed issue (or the validated discussion it +came from). See [Ticket first, PR second](#ticket-first-pr-second). + +### Run it locally + +**If you can't run it, you can't submit it.** Build the affected components and +exercise the change on a real setup — see [Build and start](#build-and-start). +"CI will tell me" is not acceptable for a VPN agent that runs as root on other +people's machines. + +### Green CI, and answer the bots + +We do not start reviewing while CI is red. Get the pipeline green first — a +failing build, lint, or test means the PR is not ready for review. + +Alongside the test workflows, your PR is reviewed by CodeRabbit and scanned by +SonarCloud, Snyk, and Codecov. Read what they report and either fix it or reply +with why it does not apply; please do not resolve the threads without a +response. They are not always right — this codebase has privileged, +platform-specific, and concurrency-heavy paths that static analysis reads poorly +— so push back when a finding is wrong rather than changing correct code to +silence it. Security and dependency findings are the exception: treat those as +real until shown otherwise. Do not edit workflows, thresholds, or scanner +configuration to make a check pass. + +### One PR, one purpose + +Bug fix, refactor, feature: separate PRs. Mixed PRs are slow to review, hard to +revert, and may be closed with a request to split them. + +### Keep it small + +Size is the strongest predictor of how long a PR waits for review. Aim for under +roughly 400 changed lines across under 20 files. Past about 1000 lines or 50 +files, expect to be asked to split the change — and large PRs from outside the +core team may be blocked until the scope has been agreed in a ticket. This is +not only about reviewer time: NetBird's agent runs as root on other people's +machines, and a sprawling diff cannot be reviewed with the care that deserves. + +Measure by hand-written code, excluding generated output, `go.sum`, and +fixtures. If a change genuinely cannot be small — a protocol migration, a +cross-component rename — agree the split in the issue before you start, and land +it as a series of PRs that each build and make sense on their own. + +### Avoid force-pushing during review + +Once a PR is open, push new commits instead of rewriting history. A force-push +detaches existing review comments from their lines, throws away the +"changes since your last review" diff, and loses the CI history that showed +which commit broke what. Since we squash on merge, there is nothing to gain from +a tidy branch history. + +Force-pushing is sometimes unavoidable — rebasing to clear a real conflict, or +removing a secret or large binary committed by mistake. When that happens, leave +a comment on the PR so reviewers know their anchors moved. + +### Quality checks + +Run these from the repository root before pushing: + +```shell +go fmt ./... +make lint # golangci-lint on files changed against origin/main +make lint-all # full-repository lint, matches CI +make test-unit # host-safe unit tests +``` + +`make setup-hooks` wires `make lint` into a pre-push hook so the fast lint runs +automatically. If your change touches privileged paths (firewall, routing, +interface management), also run `make test-privileged`, which executes the +`privileged`-tagged suite inside a Docker container with `NET_ADMIN`. + +### Code standards + - Keep functions as simple as possible, with a single purpose - Use private functions and constants where possible - Comment on any new public functions - Add unit tests for any new public function +- Comment the **why**, not the **what** — explain non-obvious decisions, invariants, and constraints, not the line below +- Keep comments within 90 characters per line and roughly 250 characters per comment; when a block needs more explanation than that, extract a named function instead of writing a longer comment (see [AGENTS.md](AGENTS.md#length-budget)) + +### PR title and commits + +PR titles must start with a bracketed tag, enforced by +[pr-title-check.yml](/.github/workflows/pr-title-check.yml): + +```text +[client] Authorize daemon IPC callers by their local identity +[management,client] Add MDM policy support +``` + +Use a comma-separated list inside a single pair of brackets when a change spans +components. The `allowedTags` array in +[pr-title-check.yml](/.github/workflows/pr-title-check.yml) is the source of +truth — at the time of writing it accepts `management`, `client`, `signal`, +`proxy`, `relay`, `misc`, `infrastructure`, `self-hosted`, and `doc`. + +Commit subjects follow the same convention — keep them short and put the +reasoning in the body, why before what, with no bullet list of files changed. + +Keep the PR description itself under 1000 words on top of the template text. +Reviewers read the diff; the description explains what the diff cannot. > When pushing fixes to the PR comments, please push as separate commits; we will squash the PR before merging, so there is no need to squash it before pushing it, and we are more than okay with 10-100 commits in a single PR. This helps review the fixes to the requested changes. +### Documentation + +User-facing changes need a matching PR in +[netbirdio/docs](https://github.com/netbirdio/docs); link it in the PR +description, or state why documentation is not needed. + +## When we close a PR + +We would rather redirect early than let a PR sit. We may close one if: + +- It changes behavior with no linked issue, or the approach was never agreed with a maintainer +- The change was clearly never run or tested locally +- CI has been red without a response +- It mixes unrelated purposes, or the purpose is not clear +- It is far too large to review and the scope was never agreed in a ticket +- The author cannot answer questions about their own change — including PRs that read as unreviewed model output, where review turns into a relay between the maintainer and an LLM. Tooling is fine; unreviewed output is not, you are responsible for the code you sign your name to +- There has been no activity for 14 days after we requested changes + +A closed PR is not a rejected idea. Take it back to the +[discussion](https://github.com/netbirdio/netbird/discussions), settle the +approach, and reopen the work from there. + ## Other project repositories NetBird project is composed of 3 main repositories: From feecb993f42163475b2673e974bb90713c2af195 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Sat, 1 Aug 2026 02:12:56 +0900 Subject: [PATCH 31/32] [client] Restrict debug bundle log path and upload destinations (#6975) --- client/android/client.go | 2 +- client/cmd/debug.go | 13 +- client/configs/configs.go | 5 + client/internal/debug/debug.go | 70 +++++--- client/internal/debug/debug_ios.go | 2 +- client/internal/debug/debug_logfiles_test.go | 2 +- client/internal/debug/uilog_test.go | 64 +++++++ client/internal/debug/upload.go | 87 +++++++++- client/internal/debug/upload_test.go | 47 +++++- client/internal/ipcauth/ownedfile.go | 63 +++++++ client/internal/ipcauth/ownedfile_test.go | 64 +++++++ client/internal/ipcauth/ownedfile_unix.go | 35 ++++ .../internal/ipcauth/ownedfile_unix_test.go | 57 +++++++ client/internal/ipcauth/ownedfile_windows.go | 59 +++++++ .../ipcauth/ownedfile_windows_test.go | 78 +++++++++ client/ios/NetBirdSDK/client.go | 2 +- client/jobexec/executor.go | 2 +- client/proto/daemon.pb.go | 32 ++-- client/proto/daemon.proto | 4 + client/server/debug.go | 92 ++++++++-- client/server/debug_gate.go | 99 +++++++++++ client/server/debug_gate_test.go | 157 ++++++++++++++++++ client/server/server.go | 3 + client/ui/uilogpath.go | 10 +- 24 files changed, 976 insertions(+), 73 deletions(-) create mode 100644 client/internal/debug/uilog_test.go create mode 100644 client/internal/ipcauth/ownedfile.go create mode 100644 client/internal/ipcauth/ownedfile_test.go create mode 100644 client/internal/ipcauth/ownedfile_unix.go create mode 100644 client/internal/ipcauth/ownedfile_unix_test.go create mode 100644 client/internal/ipcauth/ownedfile_windows.go create mode 100644 client/internal/ipcauth/ownedfile_windows_test.go create mode 100644 client/server/debug_gate.go create mode 100644 client/server/debug_gate_test.go diff --git a/client/android/client.go b/client/android/client.go index 1a8dd7d09..501d7f77c 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -301,7 +301,7 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() - key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path) + key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path, false) if err != nil { return "", fmt.Errorf("upload debug bundle: %w", err) } diff --git a/client/cmd/debug.go b/client/cmd/debug.go index 57e75f663..7ddc3afc4 100644 --- a/client/cmd/debug.go +++ b/client/cmd/debug.go @@ -29,8 +29,9 @@ const errCloseConnection = "Failed to close connection: %v" var ( logFileCount uint32 systemInfoFlag bool - uploadBundleFlag bool - uploadBundleURLFlag string + uploadBundleFlag bool + uploadBundleURLFlag string + uploadBundleInsecureFlag bool ) var debugCmd = &cobra.Command{ @@ -174,10 +175,11 @@ func debugBundle(cmd *cobra.Command, _ []string) error { } if uploadBundleFlag { request.UploadURL = uploadBundleURLFlag + request.UploadInsecure = uploadBundleInsecureFlag } resp, err := client.DebugBundle(cmd.Context(), request) if err != nil { - return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message()) + return daemonCallError("bundle debug", err) } cmd.Printf("Local file:\n%s\n", resp.GetPath()) @@ -373,10 +375,11 @@ func runForDuration(cmd *cobra.Command, args []string) error { } if uploadBundleFlag { request.UploadURL = uploadBundleURLFlag + request.UploadInsecure = uploadBundleInsecureFlag } resp, err := client.DebugBundle(cmd.Context(), request) if err != nil { - return fmt.Errorf("failed to bundle debug: %v", status.Convert(err).Message()) + return daemonCallError("bundle debug", err) } if needsRestoreUp { @@ -524,10 +527,12 @@ func init() { debugBundleCmd.Flags().BoolVarP(&systemInfoFlag, "system-info", "S", true, "Adds system information to the debug bundle") debugBundleCmd.Flags().BoolVarP(&uploadBundleFlag, "upload-bundle", "U", false, "Uploads the debug bundle to a server") debugBundleCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle") + debugBundleCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root") forCmd.Flags().Uint32VarP(&logFileCount, "log-file-count", "C", 1, "Number of rotated log files to include in debug bundle") forCmd.Flags().BoolVarP(&systemInfoFlag, "system-info", "S", true, "Adds system information to the debug bundle") forCmd.Flags().BoolVarP(&uploadBundleFlag, "upload-bundle", "U", false, "Uploads the debug bundle to a server") forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle") + forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root") forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle") } diff --git a/client/configs/configs.go b/client/configs/configs.go index 8f9c3ba28..a1ecf0feb 100644 --- a/client/configs/configs.go +++ b/client/configs/configs.go @@ -6,6 +6,11 @@ import ( "runtime" ) +// UILogFile is the file name the desktop UI writes its log to. It is defined +// here so the UI (writer), the daemon's RegisterUILog validation, and the debug +// bundle collector all share one definition. +const UILogFile = "gui-client.log" + var StateDir string func init() { diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go index 2de1023e9..0f81844f6 100644 --- a/client/internal/debug/debug.go +++ b/client/internal/debug/debug.go @@ -229,7 +229,6 @@ scutil_dns.txt (macOS only): const ( clientLogFile = "client.log" - uiLogFile = "gui-client.log" errorLogFile = "netbird.err" stdoutLogFile = "netbird.out" @@ -248,6 +247,20 @@ type MetricsExporter interface { Export(w io.Writer) error } +// LogOpener opens a log file for inclusion in the bundle. It exists so that log +// files whose path was supplied by an IPC caller can be opened under a check +// the daemon defines, instead of being opened with the daemon's privileges +// unconditionally. +type LogOpener func(path string) (*os.File, error) + +func openLogFile(path string) (*os.File, error) { + f, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + return f, nil +} + type BundleGenerator struct { anonymizer *anonymize.Anonymizer @@ -257,6 +270,7 @@ type BundleGenerator struct { syncResponse *mgmProto.SyncResponse logPath string uiLogPath string + uiLogOpener LogOpener tempDir string statePath string cpuProfile []byte @@ -285,14 +299,20 @@ type GeneratorDependencies struct { SyncResponse *mgmProto.SyncResponse LogPath string UILogPath string // Absolute path to the desktop UI's gui-client.log, reported via RegisterUILog. Empty if no UI registered one. - TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used. - StatePath string // Path to the state file. If empty, the ServiceManager default path is used. - CPUProfile []byte - CapturePath string - RefreshStatus func() - ClientMetrics MetricsExporter - DaemonVersion string - CliVersion string + // UILogOpener opens the UI log and its rotated siblings. The path comes from + // a local IPC caller, so the daemon must not open it with plain os.Open: the + // opener is where the caller's right to that file is enforced. Defaults to + // os.Open, which is only correct where the path is not caller-supplied + // (mobile). + UILogOpener LogOpener + TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used. + StatePath string // Path to the state file. If empty, the ServiceManager default path is used. + CPUProfile []byte + CapturePath string + RefreshStatus func() + ClientMetrics MetricsExporter + DaemonVersion string + CliVersion string } func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGenerator { @@ -302,6 +322,11 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen logFileCount = 1 } + uiLogOpener := deps.UILogOpener + if uiLogOpener == nil { + uiLogOpener = openLogFile + } + return &BundleGenerator{ anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()), @@ -310,6 +335,7 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen syncResponse: deps.SyncResponse, logPath: deps.LogPath, uiLogPath: deps.UILogPath, + uiLogOpener: uiLogOpener, tempDir: deps.TempDir, statePath: deps.StatePath, cpuProfile: deps.CPUProfile, @@ -996,11 +1022,11 @@ func (g *BundleGenerator) addLogfile() error { logDir := filepath.Dir(g.logPath) - if err := g.addSingleLogfile(g.logPath, clientLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, g.logPath, clientLogFile); err != nil { return fmt.Errorf("add client log file to zip: %w", err) } - g.addRotatedLogFiles(logDir, clientLogPrefix) + g.addRotatedLogFiles(openLogFile, logDir, clientLogPrefix) stdErrLogPath := filepath.Join(logDir, errorLogFile) stdoutLogPath := filepath.Join(logDir, stdoutLogFile) @@ -1009,11 +1035,11 @@ func (g *BundleGenerator) addLogfile() error { stdoutLogPath = darwinStdoutLogPath } - if err := g.addSingleLogfile(stdErrLogPath, errorLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, stdErrLogPath, errorLogFile); err != nil { log.Warnf("Failed to add %s to zip: %v", errorLogFile, err) } - if err := g.addSingleLogfile(stdoutLogPath, stdoutLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, stdoutLogPath, stdoutLogFile); err != nil { log.Warnf("Failed to add %s to zip: %v", stdoutLogFile, err) } @@ -1030,18 +1056,18 @@ func (g *BundleGenerator) addUILog() error { return nil } - if err := g.addSingleLogfile(g.uiLogPath, uiLogFile); err != nil { + if err := g.addSingleLogfile(g.uiLogOpener, g.uiLogPath, configs.UILogFile); err != nil { return fmt.Errorf("add UI log file to zip: %w", err) } - g.addRotatedLogFiles(filepath.Dir(g.uiLogPath), uiLogPrefix) + g.addRotatedLogFiles(g.uiLogOpener, filepath.Dir(g.uiLogPath), uiLogPrefix) return nil } // addSingleLogfile adds a single log file to the archive -func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error { - logFile, err := os.Open(logPath) +func (g *BundleGenerator) addSingleLogfile(open LogOpener, logPath, targetName string) error { + logFile, err := open(logPath) if err != nil { return fmt.Errorf("open log file %s: %w", targetName, err) } @@ -1066,8 +1092,8 @@ func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error { } // addSingleLogFileGz adds a single gzipped log file to the archive -func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error { - f, err := os.Open(logPath) +func (g *BundleGenerator) addSingleLogFileGz(open LogOpener, logPath, targetName string) error { + f, err := open(logPath) if err != nil { return fmt.Errorf("open gz log file %s: %w", targetName, err) } @@ -1114,7 +1140,7 @@ func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error { // addRotatedLogFiles adds rotated log files to the bundle based on logFileCount. // prefix is the base log name without extension (e.g. "client", "gui-client"); // the glob matches both files rotated by us and by logrotate on linux. -func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) { +func (g *BundleGenerator) addRotatedLogFiles(open LogOpener, logDir, prefix string) { if g.logFileCount == 0 { return } @@ -1154,9 +1180,9 @@ func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) { for i := 0; i < maxFiles; i++ { name := filepath.Base(files[i]) if strings.HasSuffix(name, ".gz") { - err = g.addSingleLogFileGz(files[i], name) + err = g.addSingleLogFileGz(open, files[i], name) } else { - err = g.addSingleLogfile(files[i], name) + err = g.addSingleLogfile(open, files[i], name) } if err != nil { log.Warnf("failed to add rotated log %s: %v", name, err) diff --git a/client/internal/debug/debug_ios.go b/client/internal/debug/debug_ios.go index a07c23dbd..001d64241 100644 --- a/client/internal/debug/debug_ios.go +++ b/client/internal/debug/debug_ios.go @@ -27,7 +27,7 @@ func (g *BundleGenerator) addPlatformLog() error { } swiftLogPath := filepath.Join(filepath.Dir(g.logPath), swiftLogFile) - if err := g.addSingleLogfile(swiftLogPath, swiftLogFile); err != nil { + if err := g.addSingleLogfile(openLogFile, swiftLogPath, swiftLogFile); err != nil { // The Swift log is best-effort: the app may not have written it yet. log.Warnf("failed to add %s to debug bundle: %v", swiftLogFile, err) } diff --git a/client/internal/debug/debug_logfiles_test.go b/client/internal/debug/debug_logfiles_test.go index 3749b3e9b..31420711f 100644 --- a/client/internal/debug/debug_logfiles_test.go +++ b/client/internal/debug/debug_logfiles_test.go @@ -97,7 +97,7 @@ func runAddRotatedLogFilesPrefix(t *testing.T, dir, prefix string, logFileCount archive: zip.NewWriter(&buf), logFileCount: logFileCount, } - g.addRotatedLogFiles(dir, prefix) + g.addRotatedLogFiles(openLogFile, dir, prefix) require.NoError(t, g.archive.Close()) zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len())) diff --git a/client/internal/debug/uilog_test.go b/client/internal/debug/uilog_test.go new file mode 100644 index 000000000..103e98c6f --- /dev/null +++ b/client/internal/debug/uilog_test.go @@ -0,0 +1,64 @@ +package debug + +import ( + "archive/zip" + "bytes" + "fmt" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/configs" +) + +// bundleEntries generates a bundle with the given generator and returns the +// set of entry names in the resulting archive. +func bundleEntries(t *testing.T, g *BundleGenerator) map[string]struct{} { + t.Helper() + + path, err := g.Generate() + require.NoError(t, err) + t.Cleanup(func() { _ = os.Remove(path) }) + + data, err := os.ReadFile(path) + require.NoError(t, err) + + zr, err := zip.NewReader(bytes.NewReader(data), int64(len(data))) + require.NoError(t, err) + + names := make(map[string]struct{}, len(zr.File)) + for _, f := range zr.File { + names[f.Name] = struct{}{} + } + return names +} + +func TestBundleIncludesUILogWhenOpenerAllows(t *testing.T) { + path := filepath.Join(t.TempDir(), configs.UILogFile) + require.NoError(t, os.WriteFile(path, []byte("gui log"), 0600)) + + g := NewBundleGenerator(GeneratorDependencies{ + UILogPath: path, + UILogOpener: openLogFile, + }, BundleConfig{}) + + require.Contains(t, bundleEntries(t, g), configs.UILogFile) +} + +// A UILogOpener that refuses (as the ownership check does for a foreign file) +// keeps the UI log out of the bundle without failing bundle generation. +func TestBundleExcludesUILogWhenOpenerRefuses(t *testing.T) { + path := filepath.Join(t.TempDir(), configs.UILogFile) + require.NoError(t, os.WriteFile(path, []byte("secret"), 0600)) + + g := NewBundleGenerator(GeneratorDependencies{ + UILogPath: path, + UILogOpener: func(string) (*os.File, error) { + return nil, fmt.Errorf("not owned by the caller") + }, + }, BundleConfig{}) + + require.NotContains(t, bundleEntries(t, g), configs.UILogFile) +} diff --git a/client/internal/debug/upload.go b/client/internal/debug/upload.go index cdf52409d..88fde6d6f 100644 --- a/client/internal/debug/upload.go +++ b/client/internal/debug/upload.go @@ -3,10 +3,12 @@ package debug import ( "context" "crypto/sha256" + "crypto/tls" "encoding/json" "fmt" "io" "net/http" + neturl "net/url" "os" "github.com/netbirdio/netbird/upload-server/types" @@ -14,20 +16,80 @@ import ( const maxBundleUploadSize = 50 * 1024 * 1024 -func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string) (key string, err error) { - response, err := getUploadURL(ctx, url, managementURL) +// requireHTTPS refuses any URL the daemon would fetch or upload to that is not +// https. The daemon runs as root and the bundle carries its logs and state, so a +// plaintext hop is a place to intercept the bundle or the presigned redirect. +// The server-side gate already enforces this for the desktop path; this also +// covers the mobile and job-runner callers that reach this package directly. +// Skipped when the caller opted into an insecure upload (self-hosted server). +func requireHTTPS(what, rawURL string) error { + parsed, err := neturl.Parse(rawURL) + if err != nil { + return fmt.Errorf("parse %s: %w", what, err) + } + if parsed.Scheme != "https" { + return fmt.Errorf("%s must use https, got scheme %q", what, parsed.Scheme) + } + return nil +} + +// uploadClient returns the HTTP client for the upload requests. The default +// client verifies TLS and refuses a redirect that would downgrade to a non-https +// hop, so a bundle can never leave over http after an https start. The insecure +// variant accepts http and untrusted certificates, and is only reachable for a +// privileged caller that passed --upload-bundle-insecure (see +// requirePrivilegeForUploadURL). +func uploadClient(insecure bool) *http.Client { + if !insecure { + return &http.Client{CheckRedirect: rejectInsecureRedirect} + } + return &http.Client{ + Transport: &http.Transport{ + //nolint:gosec // opt-in, privileged, self-hosted upload servers + TLSClientConfig: &tls.Config{InsecureSkipVerify: true, MinVersion: tls.VersionTLS12}, + }, + } +} + +// rejectInsecureRedirect refuses a redirect to a non-https target and keeps the +// standard library's 10-hop limit that a custom CheckRedirect would otherwise +// disable. +func rejectInsecureRedirect(req *http.Request, via []*http.Request) error { + if req.URL.Scheme != "https" { + return fmt.Errorf("refusing redirect to non-https URL %s", req.URL.Redacted()) + } + if len(via) >= 10 { + return fmt.Errorf("stopped after 10 redirects") + } + return nil +} + +func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string, insecure bool) (key string, err error) { + if !insecure { + if err := requireHTTPS("upload service URL", url); err != nil { + return "", err + } + } + + response, err := getUploadURL(ctx, url, managementURL, insecure) if err != nil { return "", err } - err = upload(ctx, filePath, response) + if !insecure { + if err := requireHTTPS("upload URL from service", response.URL); err != nil { + return "", err + } + } + + err = upload(ctx, filePath, response, insecure) if err != nil { return "", err } return response.Key, nil } -func upload(ctx context.Context, filePath string, response *types.GetURLResponse) error { +func upload(ctx context.Context, filePath string, response *types.GetURLResponse, insecure bool) error { fileData, err := os.Open(filePath) if err != nil { return fmt.Errorf("open file: %w", err) @@ -52,7 +114,7 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse req.ContentLength = stat.Size() req.Header.Set("Content-Type", "application/octet-stream") - putResp, err := http.DefaultClient.Do(req) + putResp, err := uploadClient(insecure).Do(req) if err != nil { return fmt.Errorf("upload failed: %v", err) } @@ -65,16 +127,23 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse return nil } -func getUploadURL(ctx context.Context, url string, managementURL string) (*types.GetURLResponse, error) { - id := getURLHash(managementURL) - getReq, err := http.NewRequestWithContext(ctx, "GET", url+"?id="+id, nil) +func getUploadURL(ctx context.Context, serviceURL string, managementURL string, insecure bool) (*types.GetURLResponse, error) { + parsed, err := neturl.Parse(serviceURL) + if err != nil { + return nil, fmt.Errorf("parse upload service URL: %w", err) + } + q := parsed.Query() + q.Set("id", getURLHash(managementURL)) + parsed.RawQuery = q.Encode() + + getReq, err := http.NewRequestWithContext(ctx, "GET", parsed.String(), nil) if err != nil { return nil, fmt.Errorf("create GET request: %w", err) } getReq.Header.Set(types.ClientHeader, types.ClientHeaderValue) - resp, err := http.DefaultClient.Do(getReq) + resp, err := uploadClient(insecure).Do(getReq) if err != nil { return nil, fmt.Errorf("get presigned URL: %w", err) } diff --git a/client/internal/debug/upload_test.go b/client/internal/debug/upload_test.go index f224b8d3f..f3927cb81 100644 --- a/client/internal/debug/upload_test.go +++ b/client/internal/debug/upload_test.go @@ -5,6 +5,7 @@ import ( "errors" "net" "net/http" + "net/http/httptest" "os" "path/filepath" "testing" @@ -43,7 +44,7 @@ func TestUpload(t *testing.T) { fileContent := []byte("test file content") err := os.WriteFile(file, fileContent, 0640) require.NoError(t, err) - key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file) + key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file, true) require.NoError(t, err) id := getURLHash(testURL) require.Contains(t, key, id+"/") @@ -79,3 +80,47 @@ func waitForServer(t *testing.T, addr string) { } t.Fatalf("server did not start listening on %s in time", addr) } + +func TestRequireHTTPS(t *testing.T) { + require.NoError(t, requireHTTPS("upload URL", "https://upload.example/path")) + require.Error(t, requireHTTPS("upload URL", "http://upload.example/path")) + require.Error(t, requireHTTPS("upload URL", "ftp://upload.example/path")) + require.Error(t, requireHTTPS("upload URL", "://malformed")) +} + +func TestRejectInsecureRedirect(t *testing.T) { + httpsReq, err := http.NewRequest(http.MethodGet, "https://a.example/", nil) + require.NoError(t, err) + require.NoError(t, rejectInsecureRedirect(httpsReq, nil), "https redirect target must be allowed") + + httpReq, err := http.NewRequest(http.MethodGet, "http://a.example/", nil) + require.NoError(t, err) + require.Error(t, rejectInsecureRedirect(httpReq, nil), "http redirect target must be refused") + + require.Error(t, rejectInsecureRedirect(httpsReq, make([]*http.Request, 10)), "the 10-redirect limit must be enforced") +} + +// The secure client refuses to follow an https response that redirects to http, +// so a bundle can't be downgraded onto plaintext mid-flight. +func TestUploadClientRefusesHTTPSToHTTPRedirect(t *testing.T) { + plain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(plain.Close) + + secure := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, plain.URL, http.StatusFound) + })) + t.Cleanup(secure.Close) + + client := uploadClient(false) + // Trust the test server's cert without disabling verification globally. + client.Transport = secure.Client().Transport + + resp, err := client.Get(secure.URL) + if resp != nil { + _ = resp.Body.Close() + } + require.Error(t, err, "redirect from https to http must be refused") + require.Contains(t, err.Error(), "non-https") +} diff --git a/client/internal/ipcauth/ownedfile.go b/client/internal/ipcauth/ownedfile.go new file mode 100644 index 000000000..be7bf4864 --- /dev/null +++ b/client/internal/ipcauth/ownedfile.go @@ -0,0 +1,63 @@ +package ipcauth + +import ( + "fmt" + "os" +) + +// OpenOwnedFile opens path for reading on behalf of the IPC caller identified by +// id, and fails unless the opened file is a regular file that id owns. +// +// It exists for the paths a local caller hands to the daemon over the IPC. The +// daemon runs as root, so opening such a path unchecked lets any local user read +// any file through it. Ownership is the invariant that keeps the daemon from +// reading, with its own privileges, a file the caller could not read itself: a +// symlink or hard link planted at the path resolves to a file someone else owns +// and is refused. +// +// The check is made against the open descriptor rather than the path, so +// swapping the path between the check and the read cannot change the answer. +// +// A privileged caller is exempt: it can read the file directly, so refusing it +// here would protect nothing. The regular-file requirement still applies to +// everyone, since a fifo or device planted at the path is never a log file. +func OpenOwnedFile(id Identity, path string) (*os.File, error) { + f, err := openForRead(path) + if err != nil { + return nil, err + } + + if err := checkOwnership(id, f); err != nil { + if cerr := f.Close(); cerr != nil { + return nil, fmt.Errorf("%w (close: %v)", err, cerr) + } + return nil, err + } + + return f, nil +} + +func checkOwnership(id Identity, f *os.File) error { + info, err := f.Stat() + if err != nil { + return fmt.Errorf("stat %s: %w", f.Name(), err) + } + + if !info.Mode().IsRegular() { + return fmt.Errorf("%s is not a regular file", f.Name()) + } + + if IsPrivilegedCaller(id) { + return nil + } + + owned, err := fileOwnedBy(id, f) + if err != nil { + return fmt.Errorf("read owner of %s: %w", f.Name(), err) + } + if !owned { + return fmt.Errorf("%s is not owned by the caller (%s)", f.Name(), id) + } + + return nil +} diff --git a/client/internal/ipcauth/ownedfile_test.go b/client/internal/ipcauth/ownedfile_test.go new file mode 100644 index 000000000..b8178ad45 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_test.go @@ -0,0 +1,64 @@ +package ipcauth + +import ( + "io" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/require" +) + +// otherIdentity is an unprivileged caller that owns nothing the test creates. +func otherIdentity(t *testing.T) Identity { + t.Helper() + if runtime.GOOS == "windows" { + return Identity{SID: "S-1-5-21-1-2-3-1001"} + } + return Identity{UID: uint32(os.Geteuid() + 1), GID: uint32(os.Getegid() + 1)} +} + +func TestOpenOwnedFileReadsFileOwnedByCaller(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("hello"), 0600)) + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + f, err := OpenOwnedFile(id, path) + require.NoError(t, err) + t.Cleanup(func() { _ = f.Close() }) + + content, err := io.ReadAll(f) + require.NoError(t, err) + require.Equal(t, "hello", string(content)) +} + +func TestOpenOwnedFileRefusesFileOwnedByAnother(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("secret"), 0600)) + + _, err := OpenOwnedFile(otherIdentity(t), path) + require.ErrorContains(t, err, "not owned by the caller") +} + +func TestOpenOwnedFileRefusesNonRegularFile(t *testing.T) { + dir := t.TempDir() + + // The caller owns the directory, so this is the regular-file requirement + // talking, not the ownership check. + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, dir) + require.ErrorContains(t, err, "not a regular file") +} + +func TestOpenOwnedFileRefusesMissingFile(t *testing.T) { + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, filepath.Join(t.TempDir(), "absent.log")) + require.Error(t, err) +} diff --git a/client/internal/ipcauth/ownedfile_unix.go b/client/internal/ipcauth/ownedfile_unix.go new file mode 100644 index 000000000..4a6afcea7 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_unix.go @@ -0,0 +1,35 @@ +//go:build !windows + +package ipcauth + +import ( + "fmt" + "os" + "syscall" +) + +// openForRead opens a caller-supplied path without following a symlink at its +// final component and without blocking: a fifo planted at the path would +// otherwise stall the open until a writer appears, and the daemon holds a lock +// while it collects the file. +func openForRead(path string) (*os.File, error) { + f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_NONBLOCK, 0) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + return f, nil +} + +func fileOwnedBy(id Identity, f *os.File) (bool, error) { + info, err := f.Stat() + if err != nil { + return false, err + } + + stat, ok := info.Sys().(*syscall.Stat_t) + if !ok { + return false, fmt.Errorf("no owner information in %T", info.Sys()) + } + + return stat.Uid == id.UID, nil +} diff --git a/client/internal/ipcauth/ownedfile_unix_test.go b/client/internal/ipcauth/ownedfile_unix_test.go new file mode 100644 index 000000000..9e7831991 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_unix_test.go @@ -0,0 +1,57 @@ +//go:build !windows + +package ipcauth + +import ( + "errors" + "os" + "path/filepath" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// A symlink is the shape the arbitrary-read attempt takes: the caller owns the +// link, the file it points at belongs to someone else. +func TestOpenOwnedFileRefusesSymlink(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "target.log") + require.NoError(t, os.WriteFile(target, []byte("secret"), 0600)) + + link := filepath.Join(dir, "gui-client.log") + require.NoError(t, os.Symlink(target, link)) + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, link) + // O_NOFOLLOW on a symlink reports ELOOP on Linux/Darwin and EMLINK on FreeBSD. + if !errors.Is(err, syscall.ELOOP) && !errors.Is(err, syscall.EMLINK) { + t.Fatalf("symlink open: got %v, want ELOOP or EMLINK", err) + } +} + +// A fifo would block the open until a writer showed up, stalling the daemon +// while it holds its lock. +func TestOpenOwnedFileRefusesFifoWithoutBlocking(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, syscall.Mkfifo(path, 0600)) + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + done := make(chan error, 1) + go func() { + _, err := OpenOwnedFile(id, path) + done <- err + }() + + select { + case err := <-done: + require.ErrorContains(t, err, "not a regular file") + case <-time.After(5 * time.Second): + t.Fatal("opening a fifo blocked") + } +} diff --git a/client/internal/ipcauth/ownedfile_windows.go b/client/internal/ipcauth/ownedfile_windows.go new file mode 100644 index 000000000..19acaec4d --- /dev/null +++ b/client/internal/ipcauth/ownedfile_windows.go @@ -0,0 +1,59 @@ +//go:build windows + +package ipcauth + +import ( + "fmt" + "os" + + "golang.org/x/sys/windows" +) + +// openForRead opens a caller-supplied path without following a reparse point at +// it. FILE_FLAG_OPEN_REPARSE_POINT is the Windows analogue of O_NOFOLLOW: it +// opens a symlink/junction itself rather than its target, so the regular-file +// check in checkOwnership refuses a link the caller planted to redirect the +// read. FILE_FLAG_BACKUP_SEMANTICS lets a directory open too (as os.Open does), +// so a directory planted at the path is refused as non-regular rather than +// erroring here. The share mode matches os.Open so a log being written stays +// openable. +func openForRead(path string) (*os.File, error) { + p, err := windows.UTF16PtrFromString(path) + if err != nil { + return nil, fmt.Errorf("convert path %s: %w", path, err) + } + + handle, err := windows.CreateFile( + p, + windows.GENERIC_READ, + windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE, + nil, + windows.OPEN_EXISTING, + windows.FILE_FLAG_OPEN_REPARSE_POINT|windows.FILE_FLAG_BACKUP_SEMANTICS, + 0, + ) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + + return os.NewFile(uintptr(handle), path), nil +} + +// fileOwnedBy compares the file's owner SID with the caller's. Files an elevated +// process creates are owned by BUILTIN\Administrators rather than by the user, +// but such a caller is privileged and never reaches this check. +func fileOwnedBy(id Identity, f *os.File) (bool, error) { + // x/sys/windows GetSecurityInfo frees the OS buffer itself and returns a + // Go-heap copy, so there is nothing to LocalFree here. + sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION) + if err != nil { + return false, fmt.Errorf("read security info: %w", err) + } + + owner, _, err := sd.Owner() + if err != nil { + return false, fmt.Errorf("read owner: %w", err) + } + + return id.SID != "" && owner.String() == id.SID, nil +} diff --git a/client/internal/ipcauth/ownedfile_windows_test.go b/client/internal/ipcauth/ownedfile_windows_test.go new file mode 100644 index 000000000..ab68fdf39 --- /dev/null +++ b/client/internal/ipcauth/ownedfile_windows_test.go @@ -0,0 +1,78 @@ +//go:build windows + +package ipcauth + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" + "golang.org/x/sys/windows" +) + +// fileOwnerSID reads the owner SID of path the same way OpenOwnedFile does, so +// the test can construct an Identity that matches (or deliberately does not). +func fileOwnerSID(t *testing.T, path string) string { + t.Helper() + f, err := os.Open(path) + require.NoError(t, err) + t.Cleanup(func() { _ = f.Close() }) + + sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION) + require.NoError(t, err) + owner, _, err := sd.Owner() + require.NoError(t, err) + return owner.String() +} + +// The allow branch of fileOwnedBy is the SID-equality path the legitimate GUI +// flow depends on. Running elevated, a created file is owned by +// BUILTIN\Administrators; an Identity carrying that SID with Elevated=false and +// no groups is unprivileged by IsPrivileged (which reads the token, not the +// SID's RID), so this exercises the real GetSecurityInfo equality rather than +// the privileged-caller shortcut. +func TestOpenOwnedFileWindowsOwnerMatchAllows(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("hello"), 0600)) + + ownerSID := fileOwnerSID(t, path) + id := Identity{SID: ownerSID} + require.False(t, id.IsPrivileged(), "identity built from the owner SID must be unprivileged for this to test the match path") + + f, err := OpenOwnedFile(id, path) + require.NoError(t, err) + _ = f.Close() +} + +func TestOpenOwnedFileWindowsOwnerMismatchRefuses(t *testing.T) { + path := filepath.Join(t.TempDir(), "gui-client.log") + require.NoError(t, os.WriteFile(path, []byte("secret"), 0600)) + + other := Identity{SID: "S-1-5-21-9-9-9-9999"} + require.False(t, other.IsPrivileged()) + + _, err := OpenOwnedFile(other, path) + require.ErrorContains(t, err, "not owned by the caller") +} + +// FILE_FLAG_OPEN_REPARSE_POINT must make OpenOwnedFile refuse a symlink the same +// way O_NOFOLLOW does on Unix, so a planted link can't redirect the read to +// another file. Creating a symlink needs a privilege the runner may lack, so the +// test skips rather than fails when it can't. +func TestOpenOwnedFileWindowsRefusesSymlink(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "target.log") + require.NoError(t, os.WriteFile(target, []byte("secret"), 0600)) + + link := filepath.Join(dir, "gui-client.log") + if err := os.Symlink(target, link); err != nil { + t.Skipf("cannot create symlink (privilege not held?): %v", err) + } + + id, err := CurrentProcessIdentity() + require.NoError(t, err) + + _, err = OpenOwnedFile(id, link) + require.Error(t, err, "a symlink must be refused") +} diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index 9289a3910..37d3e5d99 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -262,7 +262,7 @@ func (c *Client) DebugBundle(anonymize bool) (string, error) { uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) defer cancel() - key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path) + key, err := debug.UploadDebugBundle(uploadCtx, types.DefaultBundleURL, cfg.ManagementURL.String(), path, false) if err != nil { return "", fmt.Errorf("upload debug bundle: %w", err) } diff --git a/client/jobexec/executor.go b/client/jobexec/executor.go index e29cc8840..9401acacc 100644 --- a/client/jobexec/executor.go +++ b/client/jobexec/executor.go @@ -54,7 +54,7 @@ func (e *Executor) BundleJob(ctx context.Context, debugBundleDependencies debug. } }() - key, err := debug.UploadDebugBundle(ctx, types.DefaultBundleURL, mgmURL, path) + key, err := debug.UploadDebugBundle(ctx, types.DefaultBundleURL, mgmURL, path, false) if err != nil { log.Errorf("failed to upload debug bundle: %v", err) return "", fmt.Errorf("upload debug bundle: %w", err) diff --git a/client/proto/daemon.pb.go b/client/proto/daemon.pb.go index 6fbb09958..d4deeb8ec 100644 --- a/client/proto/daemon.pb.go +++ b/client/proto/daemon.pb.go @@ -2771,14 +2771,18 @@ func (x *ForwardingRulesResponse) GetRules() []*ForwardingRule { // DebugBundler type DebugBundleRequest struct { - state protoimpl.MessageState `protogen:"open.v1"` - Anonymize bool `protobuf:"varint,1,opt,name=anonymize,proto3" json:"anonymize,omitempty"` - SystemInfo bool `protobuf:"varint,3,opt,name=systemInfo,proto3" json:"systemInfo,omitempty"` - UploadURL string `protobuf:"bytes,4,opt,name=uploadURL,proto3" json:"uploadURL,omitempty"` - LogFileCount uint32 `protobuf:"varint,5,opt,name=logFileCount,proto3" json:"logFileCount,omitempty"` - CliVersion string `protobuf:"bytes,6,opt,name=cliVersion,proto3" json:"cliVersion,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache + state protoimpl.MessageState `protogen:"open.v1"` + Anonymize bool `protobuf:"varint,1,opt,name=anonymize,proto3" json:"anonymize,omitempty"` + SystemInfo bool `protobuf:"varint,3,opt,name=systemInfo,proto3" json:"systemInfo,omitempty"` + UploadURL string `protobuf:"bytes,4,opt,name=uploadURL,proto3" json:"uploadURL,omitempty"` + LogFileCount uint32 `protobuf:"varint,5,opt,name=logFileCount,proto3" json:"logFileCount,omitempty"` + CliVersion string `protobuf:"bytes,6,opt,name=cliVersion,proto3" json:"cliVersion,omitempty"` + // uploadInsecure allows uploading to an http endpoint or one with an + // untrusted TLS certificate. Restricted to privileged callers; for + // self-hosted upload servers. + UploadInsecure bool `protobuf:"varint,7,opt,name=uploadInsecure,proto3" json:"uploadInsecure,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache } func (x *DebugBundleRequest) Reset() { @@ -2846,6 +2850,13 @@ func (x *DebugBundleRequest) GetCliVersion() string { return "" } +func (x *DebugBundleRequest) GetUploadInsecure() bool { + if x != nil { + return x.UploadInsecure + } + return false +} + type DebugBundleResponse struct { state protoimpl.MessageState `protogen:"open.v1"` Path string `protobuf:"bytes,1,opt,name=path,proto3" json:"path,omitempty"` @@ -7242,7 +7253,7 @@ const file_daemon_proto_rawDesc = "" + "\x12translatedHostname\x18\x04 \x01(\tR\x12translatedHostname\x128\n" + "\x0etranslatedPort\x18\x05 \x01(\v2\x10.daemon.PortInfoR\x0etranslatedPort\"G\n" + "\x17ForwardingRulesResponse\x12,\n" + - "\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\xb4\x01\n" + + "\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\xdc\x01\n" + "\x12DebugBundleRequest\x12\x1c\n" + "\tanonymize\x18\x01 \x01(\bR\tanonymize\x12\x1e\n" + "\n" + @@ -7252,7 +7263,8 @@ const file_daemon_proto_rawDesc = "" + "\flogFileCount\x18\x05 \x01(\rR\flogFileCount\x12\x1e\n" + "\n" + "cliVersion\x18\x06 \x01(\tR\n" + - "cliVersion\"}\n" + + "cliVersion\x12&\n" + + "\x0euploadInsecure\x18\a \x01(\bR\x0euploadInsecure\"}\n" + "\x13DebugBundleResponse\x12\x12\n" + "\x04path\x18\x01 \x01(\tR\x04path\x12 \n" + "\vuploadedKey\x18\x02 \x01(\tR\vuploadedKey\x120\n" + diff --git a/client/proto/daemon.proto b/client/proto/daemon.proto index 8d5294eb7..3c31156ec 100644 --- a/client/proto/daemon.proto +++ b/client/proto/daemon.proto @@ -536,6 +536,10 @@ message DebugBundleRequest { string uploadURL = 4; uint32 logFileCount = 5; string cliVersion = 6; + // uploadInsecure allows uploading to an http endpoint or one with an + // untrusted TLS certificate. Restricted to privileged callers; for + // self-hosted upload servers. + bool uploadInsecure = 7; } message DebugBundleResponse { diff --git a/client/server/debug.go b/client/server/debug.go index 0b6ac4b53..60a401b0e 100644 --- a/client/server/debug.go +++ b/client/server/debug.go @@ -7,18 +7,62 @@ import ( "context" "errors" "fmt" + "path/filepath" "runtime/pprof" + "strings" + "time" log "github.com/sirupsen/logrus" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" "github.com/netbirdio/netbird/client/internal/debug" + "github.com/netbirdio/netbird/client/internal/ipcauth" "github.com/netbirdio/netbird/client/proto" mgmProto "github.com/netbirdio/netbird/shared/management/proto" "github.com/netbirdio/netbird/version" ) // DebugBundle creates a debug bundle and returns the location. -func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) (resp *proto.DebugBundleResponse, err error) { +func (s *Server) DebugBundle(callerCtx context.Context, req *proto.DebugBundleRequest) (resp *proto.DebugBundleResponse, err error) { + if err := requirePrivilegeForUploadURL(callerCtx, req.GetUploadURL(), req.GetUploadInsecure()); err != nil { + return nil, err + } + + // The UI log is opened as whoever asked for this bundle, so a caller only + // collects a log it owns (privileged callers excepted). ok is false on a + // socket that carries no identity, which skips the UI log. + callerID, callerIdentified := ipcauth.CallerIdentity(callerCtx) + + path, managementURL, err := s.generateDebugBundle(req, uiLogOpener(callerID, callerIdentified)) + if err != nil { + return nil, err + } + + if req.GetUploadURL() == "" { + return &proto.DebugBundleResponse{Path: path}, nil + } + + // The upload runs without s.mutex held: it does network I/O to a possibly + // slow destination and must not block the other RPCs that take the lock. The + // bounded context is a backstop against a hung connection. + uploadCtx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() + key, err := debug.UploadDebugBundle(uploadCtx, req.GetUploadURL(), managementURL, path, req.GetUploadInsecure()) + if err != nil { + log.Errorf("failed to upload debug bundle to %s: %v", req.GetUploadURL(), err) + return &proto.DebugBundleResponse{Path: path, UploadFailureReason: err.Error()}, nil + } + + log.Infof("debug bundle uploaded to %s with key %s", req.GetUploadURL(), key) + + return &proto.DebugBundleResponse{Path: path, UploadedKey: key}, nil +} + +// generateDebugBundle builds the bundle under s.mutex and returns its path plus +// the management URL captured under the lock, so the caller can run the upload +// without holding the lock. +func (s *Server) generateDebugBundle(req *proto.DebugBundleRequest, uiOpener debug.LogOpener) (path string, managementURL string, err error) { s.mutex.Lock() defer s.mutex.Unlock() @@ -68,6 +112,7 @@ func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) ( SyncResponse: syncResponse, LogPath: s.logFile, UILogPath: s.uiLogPath, + UILogOpener: uiOpener, CPUProfile: cpuProfileData, CapturePath: capturePath, RefreshStatus: refreshStatus, @@ -82,23 +127,16 @@ func (s *Server) DebugBundle(_ context.Context, req *proto.DebugBundleRequest) ( }, ) - path, err := bundleGenerator.Generate() + path, err = bundleGenerator.Generate() if err != nil { - return nil, fmt.Errorf("generate debug bundle: %w", err) + return "", "", fmt.Errorf("generate debug bundle: %w", err) } - if req.GetUploadURL() == "" { - return &proto.DebugBundleResponse{Path: path}, nil - } - key, err := debug.UploadDebugBundle(context.Background(), req.GetUploadURL(), s.config.ManagementURL.String(), path) - if err != nil { - log.Errorf("failed to upload debug bundle to %s: %v", req.GetUploadURL(), err) - return &proto.DebugBundleResponse{Path: path, UploadFailureReason: err.Error()}, nil + if s.config != nil && s.config.ManagementURL != nil { + managementURL = s.config.ManagementURL.String() } - log.Infof("debug bundle uploaded to %s with key %s", req.GetUploadURL(), key) - - return &proto.DebugBundleResponse{Path: path, UploadedKey: key}, nil + return path, managementURL, nil } // GetLogLevel gets the current logging level for the server. @@ -138,12 +176,34 @@ func (s *Server) SetLogLevel(_ context.Context, req *proto.SetLogLevelRequest) ( // RegisterUILog records the desktop UI's absolute log path so DebugBundle can // collect the GUI log. The daemon runs as root and can't resolve the user's // config dir, so the UI reports it. Last-writer-wins (one UI per socket). -func (s *Server) RegisterUILog(_ context.Context, req *proto.RegisterUILogRequest) (*proto.RegisterUILogResponse, error) { +// +// The path arrives over an IPC any local user can reach and is later opened by +// a root daemon, so it is constrained to the file name the UI writes and to a +// local absolute path. Authorization happens when DebugBundle opens it: the +// bundle refuses a file its requester does not own. A caller the daemon cannot +// identify cannot register a path at all. +func (s *Server) RegisterUILog(callerCtx context.Context, req *proto.RegisterUILogRequest) (*proto.RegisterUILogResponse, error) { + if _, ok := ipcauth.CallerIdentity(callerCtx); !ok { + return nil, gstatus.Error(codes.PermissionDenied, + "registering a UI log path requires a control channel that carries the caller's identity") + } + + path := filepath.Clean(req.GetPath()) + if !filepath.IsAbs(path) || filepath.Base(path) != uiLogFileName { + return nil, gstatus.Errorf(codes.InvalidArgument, "UI log path must be an absolute path ending in %s", uiLogFileName) + } + // filepath.IsAbs accepts a Windows UNC path (\\host\share\...) and a device + // path (\\.\, \\?\); opening one would make the root daemon reach a remote + // or device namespace. Require a plain local path. + if strings.HasPrefix(path, `\\`) { + return nil, gstatus.Error(codes.InvalidArgument, "UI log path must be a local path, not a UNC or device path") + } + s.mutex.Lock() defer s.mutex.Unlock() - s.uiLogPath = req.GetPath() - log.Infof("registered UI log path: %s", s.uiLogPath) + s.uiLogPath = path + log.Infof("registered UI log path %s", s.uiLogPath) return &proto.RegisterUILogResponse{}, nil } diff --git a/client/server/debug_gate.go b/client/server/debug_gate.go new file mode 100644 index 000000000..983a13aaf --- /dev/null +++ b/client/server/debug_gate.go @@ -0,0 +1,99 @@ +//go:build !android && !ios + +package server + +import ( + "context" + "fmt" + "net/url" + "os" + "strings" + + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/configs" + "github.com/netbirdio/netbird/client/internal/debug" + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/upload-server/types" +) + +// uiLogFileName is the only file name the daemon accepts as a UI log path. The +// UI (writer), this validation, and the bundle collector all read it from +// configs so they cannot drift. +const uiLogFileName = configs.UILogFile + +// uiLogOpener opens the registered UI log, and its rotated siblings, on behalf +// of the caller requesting the bundle: OpenOwnedFile then collects the log only +// when that caller owns it (or is privileged). identified is false on a socket +// that carries no caller identity, in which case nothing is opened. +func uiLogOpener(id ipcauth.Identity, identified bool) debug.LogOpener { + return func(path string) (*os.File, error) { + if !identified { + return nil, fmt.Errorf("bundle requester has no verified identity") + } + return ipcauth.OpenOwnedFile(id, path) + } +} + +// requirePrivilegeForUploadURL restricts where the daemon may send a debug +// bundle. The bundle holds the daemon's own logs and state, and the daemon +// fetches the upload URL itself, so an unrestricted endpoint turns the daemon +// into both an exfiltration channel and a request forwarder that reaches +// services only it can talk to. +// +// The upload service NetBird publishes is open to any caller, since that is what +// the CLI and the desktop UI use. Any other endpoint, self-hosted upload servers +// included, requires a privileged caller. Plaintext is refused for everyone: the +// daemon fetches the URL and then PUTs the bundle to whatever that fetch returns, +// so an http hop is a place to intercept the bundle or the redirect. +// +// insecure relaxes transport security (http, or an untrusted TLS certificate) +// for a self-hosted server. It weakens a root-privileged upload, so it is +// refused for an unprivileged caller regardless of the host. +func requirePrivilegeForUploadURL(ctx context.Context, rawURL string, insecure bool) error { + if rawURL == "" { + return nil + } + + parsed, err := url.Parse(rawURL) + if err != nil { + return gstatus.Errorf(codes.InvalidArgument, "parse upload URL: %v", err) + } + + // --insecure relaxes https to http or an untrusted certificate; it does not + // widen the URL to arbitrary schemes, so a host and http/https are required + // before the insecure branch takes over. + if parsed.Host == "" || (parsed.Scheme != "https" && parsed.Scheme != "http") { + return gstatus.Errorf(codes.InvalidArgument, "upload URL must be http or https with a host") + } + + if insecure { + return denyPrivileged(ctx, + "uploading a debug bundle without transport security (--upload-bundle-insecure)", + ipcauth.ElevatedCommand("netbird debug bundle -U --upload-bundle-insecure --upload-bundle-url ")) + } + + if parsed.Scheme != "https" { + return gstatus.Errorf(codes.InvalidArgument, "upload URL must use https, got scheme %q", parsed.Scheme) + } + + if isDefaultUploadService(parsed) { + return nil + } + + return denyPrivileged(ctx, + "uploading a debug bundle to an upload service other than the default one", + ipcauth.ElevatedCommand("netbird debug bundle -U --upload-bundle-url ")) +} + +// isDefaultUploadService reports whether the URL points at the upload service +// NetBird runs. Only the host is compared: the service's path may differ between +// releases, and the host is what decides who receives the bundle. +func isDefaultUploadService(parsed *url.URL) bool { + defaultURL, err := url.Parse(types.DefaultBundleURL) + if err != nil { + return false + } + return parsed.Scheme == defaultURL.Scheme && strings.EqualFold(parsed.Host, defaultURL.Host) +} diff --git a/client/server/debug_gate_test.go b/client/server/debug_gate_test.go new file mode 100644 index 000000000..e958fc581 --- /dev/null +++ b/client/server/debug_gate_test.go @@ -0,0 +1,157 @@ +//go:build !android && !ios + +package server + +import ( + "os" + "path/filepath" + "runtime" + "testing" + + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" + + "github.com/netbirdio/netbird/client/internal/ipcauth" + "github.com/netbirdio/netbird/client/proto" + "github.com/netbirdio/netbird/upload-server/types" +) + +func TestRegisterUILogRefusesUnidentifiedCaller(t *testing.T) { + s := &Server{} + + _, err := s.RegisterUILog(noIdentityCtx(), &proto.RegisterUILogRequest{ + Path: filepath.Join(t.TempDir(), uiLogFileName), + }) + + if gstatus.Code(err) != codes.PermissionDenied { + t.Fatalf("code = %v, want PermissionDenied", gstatus.Code(err)) + } +} + +func TestRegisterUILogRefusesForeignPath(t *testing.T) { + secret := "/etc/shadow" + if runtime.GOOS == "windows" { + secret = `C:\Windows\System32\config\SAM` + } + + tests := []struct { + name string + path string + }{ + {"empty", ""}, + {"relative", filepath.Join("netbird", uiLogFileName)}, + {"another file", secret}, + {"directory of the log", t.TempDir()}, + {"unc path", `\\attacker\share\` + uiLogFileName}, + {"device path", `\\.\C:\` + uiLogFileName}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + s := &Server{} + + _, err := s.RegisterUILog(userCtx(), &proto.RegisterUILogRequest{Path: tc.path}) + + if gstatus.Code(err) != codes.InvalidArgument { + t.Fatalf("code = %v, want InvalidArgument", gstatus.Code(err)) + } + if s.uiLogPath != "" { + t.Fatalf("path %q was recorded despite the refusal", s.uiLogPath) + } + }) + } +} + +func TestRegisterUILogRecordsPath(t *testing.T) { + s := &Server{} + path := filepath.Join(t.TempDir(), uiLogFileName) + + if _, err := s.RegisterUILog(userCtx(), &proto.RegisterUILogRequest{Path: path}); err != nil { + t.Fatalf("register: %v", err) + } + + if s.uiLogPath != path { + t.Fatalf("path = %q, want %q", s.uiLogPath, path) + } +} + +// The UI log is opened as the bundle requester, so a second local user cannot +// collect a log they do not own, and an unidentified requester collects nothing. +func TestUILogOpenerBindsToRequester(t *testing.T) { + path := filepath.Join(t.TempDir(), uiLogFileName) + if err := os.WriteFile(path, []byte("log line"), 0600); err != nil { + t.Fatalf("write log: %v", err) + } + + // A different unprivileged user than the file's owner: refused. + if _, err := uiLogOpener(unprivilegedIdentity(), true)(path); err == nil { + t.Fatal("expected a file the requester does not own to be refused") + } + + // No verified identity: refused. + if _, err := uiLogOpener(ipcauth.Identity{}, false)(path); err == nil { + t.Fatal("expected an unidentified requester to be refused") + } + + // The requester that owns the file: allowed. The test process created it, so + // its own identity is the owner (and a privileged runner is exempt anyway). + owner, err := ipcauth.CurrentProcessIdentity() + if err != nil { + t.Fatalf("current identity: %v", err) + } + f, err := uiLogOpener(owner, true)(path) + if err != nil { + t.Fatalf("expected the owning requester to be allowed, got %v", err) + } + _ = f.Close() +} + +func TestRequirePrivilegeForUploadURL(t *testing.T) { + tests := []struct { + name string + url string + insecure bool + unprivOK bool + invalid bool + rootAlso bool + }{ + {name: "no upload", url: "", unprivOK: true}, + {name: "default service", url: types.DefaultBundleURL, unprivOK: true}, + {name: "default service, other path", url: "https://upload.debug.netbird.io/other", unprivOK: true}, + {name: "loopback exfiltration endpoint", url: "https://127.0.0.1:8080/upload-url", rootAlso: true}, + {name: "custom upload service", url: "https://attacker.example/upload-url", rootAlso: true}, + {name: "plaintext default host", url: "http://upload.debug.netbird.io/upload-url", invalid: true}, + {name: "plaintext custom host", url: "http://attacker.example/upload-url", invalid: true}, + {name: "unsupported scheme", url: "file:///etc/shadow", invalid: true}, + // insecure relaxes transport security; privileged only, whatever the host. + {name: "insecure http custom", url: "http://selfhosted.local/upload-url", insecure: true, rootAlso: true}, + {name: "insecure https custom", url: "https://selfhosted.local/upload-url", insecure: true, rootAlso: true}, + {name: "insecure default host", url: types.DefaultBundleURL, insecure: true, rootAlso: true}, + // --insecure must not widen the URL to non-http(s) schemes or a hostless URL. + {name: "insecure file scheme", url: "file:///etc/shadow", insecure: true, invalid: true}, + {name: "insecure hostless", url: "https:///upload-url", insecure: true, invalid: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + err := requirePrivilegeForUploadURL(userCtx(), tc.url, tc.insecure) + + switch { + case tc.invalid: + if gstatus.Code(err) != codes.InvalidArgument { + t.Fatalf("code = %v, want InvalidArgument", gstatus.Code(err)) + } + return + case tc.unprivOK: + assertAllowed(t, err) + return + default: + assertDenied(t, err) + } + + if tc.rootAlso { + assertAllowed(t, requirePrivilegeForUploadURL(rootCtx(), tc.url, tc.insecure)) + } + }) + } +} diff --git a/client/server/server.go b/client/server/server.go index 6e22a76a9..aaab5cc02 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -72,6 +72,9 @@ type Server struct { // RegisterUILog. Guarded by mutex. Consumed by DebugBundle so the bundle // can collect the GUI log even though the daemon runs as root and can't // resolve the user's config dir. Last-writer-wins (one UI per socket). + // DebugBundle opens it on behalf of the bundle requester and refuses a file + // that caller does not own, so a local user cannot read another user's log + // or a root-only file through it. uiLogPath string oauthAuthFlow oauthAuthFlow diff --git a/client/ui/uilogpath.go b/client/ui/uilogpath.go index 96fbb9637..6fa400e01 100644 --- a/client/ui/uilogpath.go +++ b/client/ui/uilogpath.go @@ -8,21 +8,19 @@ import ( log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/configs" "github.com/netbirdio/netbird/client/ui/guilog" ) -// uiLogFileName must stay in sync with the daemon's "gui-client*.log.*" glob -// for rotated siblings (addUILog in client/internal/debug). -const uiLogFileName = "gui-client.log" - // uiLogPath returns the GUI log path with native separators, since the daemon -// opens it directly for debug-bundle collection. +// opens it directly for debug-bundle collection. The file name comes from +// configs.UILogFile so the daemon validates and collects the same name. func uiLogPath() (string, error) { dir, err := os.UserConfigDir() if err != nil { return "", err } - return filepath.Join(dir, "netbird", uiLogFileName), nil + return filepath.Join(dir, "netbird", configs.UILogFile), nil } // newDebugLog builds the GUI debug log, disabled when userSetLogFile is set From 0780a806f2cc2e8a6a51782cfffe0591b7c3fa9c Mon Sep 17 00:00:00 2001 From: Misha Bragin Date: Fri, 31 Jul 2026 20:52:56 +0200 Subject: [PATCH 32/32] [management, proxy] Management-owned LLM pricing: file-backed defaults + (#6965) --- combined/cmd/config.go | 10 + combined/cmd/root.go | 52 ++ combined/config.yaml.example | 13 + e2e/agentnetwork/chat_test.go | 67 +- e2e/agentnetwork/custom_pricing_test.go | 633 ++++++++++++++++++ e2e/harness/agentnetwork.go | 21 + e2e/harness/combined.go | 11 +- e2e/harness/config.go | 39 ++ management/cmd/management.go | 24 + .../modules/agentnetwork/catalog/catalog.go | 176 +++-- .../handlers/providers_handler.go | 71 +- .../handlers/providers_handler_test.go | 53 ++ .../modules/agentnetwork/pricing/defaults.go | 156 +++++ .../pricing/defaults_llm_pricing.example.yaml | 494 +++++++------- .../agentnetwork/pricing/defaults_test.go | 148 ++++ .../agentnetwork/pricing/exampleyaml.go | 109 +++ .../modules/agentnetwork/pricing/gen.go | 20 + .../modules/agentnetwork/pricing/override.go | 249 +++++++ .../agentnetwork/pricing/override_test.go | 173 +++++ .../modules/agentnetwork/synthesizer.go | 14 +- .../agentnetwork/synthesizer_pricing.go | 131 ++++ .../agentnetwork/synthesizer_pricing_test.go | 105 +++ .../modules/agentnetwork/synthesizer_test.go | 41 +- .../modules/agentnetwork/types/provider.go | 50 +- .../modules/agentnetwork/wire_shape_test.go | 6 + management/internals/server/config/config.go | 21 + proxy/internal/llm/bedrock_model.go | 38 -- proxy/internal/llm/fixtures/pricing.yaml | 59 -- proxy/internal/llm/model.go | 21 + .../llm/pricing/defaults_coverage_test.go | 65 -- proxy/internal/llm/pricing/pricing.go | 487 ++++---------- proxy/internal/llm/pricing/pricing_other.go | 20 - proxy/internal/llm/pricing/pricing_test.go | 385 ++--------- proxy/internal/llm/pricing/pricing_unix.go | 68 -- proxy/internal/middleware/builtin/builtin.go | 7 +- .../builtin/cost_calculation_matrix_test.go | 14 +- .../middleware/builtin/cost_meter/factory.go | 82 ++- .../builtin/cost_meter/middleware.go | 68 +- .../builtin/cost_meter/middleware_test.go | 304 +++++---- .../llm_request_parser/bedrock_test.go | 17 - .../builtin/llm_request_parser/middleware.go | 39 +- .../agent_network_chain_realstack_test.go | 40 +- proxy/server.go | 6 +- shared/llm/model.go | 58 ++ .../llm/model_test.go | 13 + shared/management/http/api/openapi.yml | 37 + shared/management/http/api/types.gen.go | 21 + 47 files changed, 3196 insertions(+), 1540 deletions(-) create mode 100644 e2e/agentnetwork/custom_pricing_test.go create mode 100644 management/internals/modules/agentnetwork/handlers/providers_handler_test.go create mode 100644 management/internals/modules/agentnetwork/pricing/defaults.go rename proxy/internal/llm/pricing/defaults_pricing.yaml => management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml (63%) create mode 100644 management/internals/modules/agentnetwork/pricing/defaults_test.go create mode 100644 management/internals/modules/agentnetwork/pricing/exampleyaml.go create mode 100644 management/internals/modules/agentnetwork/pricing/gen.go create mode 100644 management/internals/modules/agentnetwork/pricing/override.go create mode 100644 management/internals/modules/agentnetwork/pricing/override_test.go create mode 100644 management/internals/modules/agentnetwork/synthesizer_pricing.go create mode 100644 management/internals/modules/agentnetwork/synthesizer_pricing_test.go delete mode 100644 proxy/internal/llm/bedrock_model.go delete mode 100644 proxy/internal/llm/fixtures/pricing.yaml create mode 100644 proxy/internal/llm/model.go delete mode 100644 proxy/internal/llm/pricing/defaults_coverage_test.go delete mode 100644 proxy/internal/llm/pricing/pricing_other.go delete mode 100644 proxy/internal/llm/pricing/pricing_unix.go create mode 100644 shared/llm/model.go rename proxy/internal/llm/bedrock_model_test.go => shared/llm/model_test.go (66%) diff --git a/combined/cmd/config.go b/combined/cmd/config.go index 7f30cd8a8..890c86876 100644 --- a/combined/cmd/config.go +++ b/combined/cmd/config.go @@ -76,6 +76,13 @@ type ServerConfig struct { SupportedSyncMessageVersions *int `yaml:"supportedSyncMessageVersions,omitempty"` PerAccountSupportedSyncMessageVersions map[string]int `yaml:"perAccountSupportedSyncMessageVersions,omitempty"` + + AgentNetwork AgentNetworkConfig `yaml:"agentNetwork"` +} + +// AgentNetworkConfig contains agent-network (LLM gateway) configuration. +type AgentNetworkConfig struct { + PricingDefaultsFile string `yaml:"pricingDefaultsFile"` } // TLSConfig contains TLS/HTTPS settings @@ -723,6 +730,9 @@ func (c *CombinedConfig) ToManagementConfig() (*nbconfig.Config, error) { EmbeddedIdP: embeddedIdP, HighestSupportedSyncMessageVersion: c.Server.SupportedSyncMessageVersions, PerAccountHighestSupportedSyncMessageVersion: c.Server.PerAccountSupportedSyncMessageVersions, + AgentNetwork: nbconfig.AgentNetwork{ + PricingDefaultsFile: c.Server.AgentNetwork.PricingDefaultsFile, + }, }, nil } diff --git a/combined/cmd/root.go b/combined/cmd/root.go index 5f2564e3a..7eac84ce5 100644 --- a/combined/cmd/root.go +++ b/combined/cmd/root.go @@ -10,6 +10,7 @@ import ( "net/http" "os" "os/signal" + "path/filepath" "strconv" "strings" "sync" @@ -24,6 +25,7 @@ import ( "google.golang.org/grpc" "github.com/netbirdio/netbird/encryption" + agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" mgmtServer "github.com/netbirdio/netbird/management/internals/server" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" "github.com/netbirdio/netbird/management/server/telemetry" @@ -288,6 +290,11 @@ func (s *serverInstances) createManagementServer(ctx context.Context, cfg *Combi return fmt.Errorf("failed to ensure encryption key: %w", err) } + if err := loadAgentNetworkPricing(ctx, mgmtConfig); err != nil { + cleanupSTUNListeners(s.stunListeners) + return fmt.Errorf("failed to load agent-network pricing defaults: %w", err) + } + LogConfigInfo(mgmtConfig) s.mgmtSrv, err = createManagementServer(cfg, mgmtConfig) @@ -622,6 +629,32 @@ func handleRelayWebSocket(w http.ResponseWriter, r *http.Request, acceptFn func( acceptFn(conn) } +// loadAgentNetworkPricing loads the management-side LLM pricing defaults +// file for the combined server and starts its periodic reloader. An +// explicitly configured PricingDefaultsFile is required to load (a typo +// must fail startup rather than silently bill with built-ins the operator +// believes they replaced); a relative path is resolved against the data +// directory so a bare filename like "pricing.yaml" lands in the datadir +// alongside the store. With no path configured, / +// is probed and may be absent (compiled-in defaults serve). +func loadAgentNetworkPricing(ctx context.Context, mgmtConfig *nbconfig.Config) error { + pricingPath := mgmtConfig.AgentNetwork.PricingDefaultsFile + required := pricingPath != "" + if !required { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(mgmtConfig.Datadir, pricingPath) + } + + log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, required) + if err := agentnetworkpricing.LoadFile(pricingPath, required); err != nil { + return err + } + agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval) + return nil +} + // logConfig prints all configuration parameters for debugging func logConfig(cfg *CombinedConfig) { log.Info("=== Configuration ===") @@ -698,6 +731,25 @@ func logManagementConfig(cfg *CombinedConfig) { log.Infof(" Relay addresses: %v", cfg.Management.Relays.Addresses) log.Infof(" Relay credentials TTL: %s", cfg.Management.Relays.CredentialsTTL) } + + logAgentNetworkConfig(cfg) +} + +func logAgentNetworkConfig(cfg *CombinedConfig) { + log.Info(" Agent Network:") + pricingPath := cfg.Server.AgentNetwork.PricingDefaultsFile + configured := pricingPath != "" + if !configured { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(cfg.Management.DataDir, pricingPath) + } + if configured { + log.Infof(" Pricing defaults file: %s", pricingPath) + } else { + log.Infof(" Pricing defaults file: %s (default, optional)", pricingPath) + } } // logEnvVars logs all NB_ environment variables that are currently set diff --git a/combined/config.yaml.example b/combined/config.yaml.example index 66bc71703..085e4344f 100644 --- a/combined/config.yaml.example +++ b/combined/config.yaml.example @@ -134,3 +134,16 @@ server: # trustedPeers: [] # CIDRs of trusted peer networks (e.g. ["100.64.0.0/10"]) # accessLogRetentionDays: 7 # Days to retain HTTP access logs. 0 (or unset) defaults to 7. Negative values disable cleanup (logs kept indefinitely). # accessLogCleanupIntervalHours: 24 # How often (in hours) to run the access-log cleanup job. 0 (or unset) is treated as "not set" and defaults to 24 hours; cleanup remains enabled. To disable cleanup, set accessLogRetentionDays to a negative value. + + # Agent network (LLM gateway) settings (optional) + # agentNetwork: + # # Path to the YAML file holding the default LLM pricing table. A relative + # # path is resolved against dataDir, so a bare filename like "pricing.yaml" + # # lands in the data directory. When empty, {dataDir}/defaults_llm_pricing.yaml + # # is probed; if no file is present the compiled-in defaults are used. + # # Schema: surface ("openai"/"anthropic"/"bedrock") -> model -> rates in USD + # # per 1k tokens (input_per_1k, output_per_1k, and the optional + # # cached_input_per_1k / cache_read_per_1k / cache_creation_per_1k). The file + # # is re-read periodically (mtime poll). An explicitly configured path that + # # fails to load fails startup; runtime reload errors keep the previous table. + # pricingDefaultsFile: "pricing.yaml" diff --git a/e2e/agentnetwork/chat_test.go b/e2e/agentnetwork/chat_test.go index 65c4d813f..ed94623fe 100644 --- a/e2e/agentnetwork/chat_test.go +++ b/e2e/agentnetwork/chat_test.go @@ -23,8 +23,13 @@ import ( type per1k struct{ in, out, read, write float64 } // publishedPer1k hardcodes the vendors' PUBLISHED rates for the models the live matrix can drive, -// keyed by the normalized model id the proxy stamps. Deliberately independent of the proxy's -// pricing table so a wrong embedded rate or a broken normalization fails the run. +// keyed by the normalized model id the proxy stamps. Deliberately independent of NetBird's own +// default pricing table so a wrong default rate or a broken normalization fails the run. +// +// These rates are also what providerRequest registers as the operator's per-model prices. Since +// management now ships operator prices to the cost meter as a per-provider-record table that is +// consulted BEFORE the surface defaults, registering the published rate is what keeps this matrix +// asserting vendor rates — and exercises the per-record path at the same time. var publishedPer1k = map[string]per1k{ "gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0}, "gpt-4o": {0.0025, 0.01, 0.00125, 0}, @@ -35,12 +40,22 @@ var publishedPer1k = map[string]per1k{ "anthropic.claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125}, "anthropic.claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375}, "anthropic.claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375}, + // Gateway-prefixed ids (Vercel AI Gateway, OpenRouter). A gateway model is not in + // NetBird's default table, so before operator pricing it could only be recorded at + // cost 0. The operator names it and prices it — at the underlying vendor's published + // rate, which is what the gateway charges through — so these rows are now priced. + "openai/gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0}, + "openai/gpt-4o": {0.0025, 0.01, 0.00125, 0}, } // rawCostVerificationSQL is the operator-facing double-check, run straight against the management // sqlite store: recompute each usage row's expected total and cache cost from its own persisted // token buckets and hardcoded published rates. OpenAI counts cached tokens as a subset of input; // Anthropic-shape providers count cache buckets additively. +// +// The rate rows must stay in sync with publishedPer1k — they are the same vendor rates the matrix +// registers as operator prices. The join is on model, so rows written by other tests in this +// package (which price their own made-up model ids) are simply not covered here. const rawCostVerificationSQL = ` WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS ( VALUES @@ -52,7 +67,9 @@ WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS ( ('kimi-k3', 0.003, 0.015, 0.0003, 0.003), ('anthropic.claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125), ('anthropic.claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375), - ('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375) + ('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375), + ('openai/gpt-4o-mini', 0.00015, 0.0006, 0.000075, 0.0), + ('openai/gpt-4o', 0.0025, 0.01, 0.00125, 0.0) ) SELECT u.provider, @@ -146,6 +163,11 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) { require.Positive(t, verified, "raw SQL check must cover at least one usage row") t.Logf("[sql] verified %d usage rows in store.db against published rates", verified) + // Gateway-prefixed model ids are absent from NetBird's default pricing table, so they are + // priced only because the operator registered and priced them on the provider record. Assert + // they are priced (not silently 0) — the join above already checked the exact figures for the + // ones this matrix drives. A gateway row at cost 0 means the per-record table never reached + // the cost meter, which is the regression this guards. gwRows, err := db.Raw(`SELECT model, (input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd) AS cost_usd FROM agent_network_request_usage WHERE model LIKE '%/%'`).Rows() @@ -155,8 +177,8 @@ func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) { var model string var cost float64 require.NoError(t, gwRows.Scan(&model, &cost), "scan gateway usage row") - t.Logf("[sql] gateway %s: stored=$%.6f (must be 0 — deliberately unpriced)", model, cost) - assert.Zerof(t, cost, "gateway-prefixed model %q must store cost 0, never a guessed rate", model) + t.Logf("[sql] gateway %s: stored=$%.6f (priced from the operator's per-record rate)", model, cost) + assert.Positivef(t, cost, "gateway-prefixed model %q is priced on the provider record, so its cost must be > 0", model) } require.NoError(t, gwRows.Err(), "iterate gateway usage rows") } @@ -177,10 +199,6 @@ func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAc rates, known := publishedPer1k[model] if !known { - if strings.Contains(model, "/") { - assert.Zerof(t, row.CostUsd, "gateway-prefixed model %q is not priced so the cost meter must skip (cost 0)", model) - return - } t.Logf("[cost] %s: no published rate on file for model %q (env-overridden?); skipping cost validation", pc.name, model) return } @@ -337,8 +355,17 @@ func availableProviders() []providerCase { } // providerRequest builds a create request for a matrix provider: enabled, with -// a uniquely-priced model for body-routed providers and none for the -// path-routed Vertex (whose model lives in the request path). +// its model registered at the vendor's published rates for body-routed +// providers, and no models for the path-routed Vertex (whose model lives in the +// request path, so it prices from the defaults table management ships). +// +// The registered rates matter: management synthesizes them into the cost +// meter's per-provider-record table, which is consulted before the surface +// defaults, so these are the rates the proxy actually bills with. Registering +// the published rate keeps the cost assertions vendor-anchored while covering +// the operator-pricing path. A model with no published rate on file (an +// env-overridden Bedrock profile) falls back to a nominal rate, and +// validateAccessLogCost skips its cost check. func providerRequest(pc providerCase) api.AgentNetworkProviderRequest { req := api.AgentNetworkProviderRequest{ Name: pc.name, @@ -356,9 +383,23 @@ func providerRequest(pc providerCase) api.AgentNetworkProviderRequest { if pc.kind == harness.WireBedrock { modelID = catalogModel(pc) } - req.Models = &[]api.AgentNetworkProviderModel{ - {Id: modelID, InputPer1k: 0.001, OutputPer1k: 0.002}, + model := api.AgentNetworkProviderModel{Id: modelID, InputPer1k: 0.001, OutputPer1k: 0.002} + if rates, known := publishedPer1k[catalogModel(pc)]; known { + model.InputPer1k = rates.in + model.OutputPer1k = rates.out + // Pin the cache rates too, rather than letting them inherit from the + // defaults table: a gateway-prefixed id has no default entry to + // inherit from, and an unset rate bills that bucket at the input + // rate, which would not match the published-rate recompute. + if rates.read > 0 { + model.CachedInputPer1k = ptr(rates.read) // OpenAI shape + model.CacheReadPer1k = ptr(rates.read) // Anthropic / Bedrock shape + } + if rates.write > 0 { + model.CacheCreationPer1k = ptr(rates.write) + } } + req.Models = &[]api.AgentNetworkProviderModel{model} } return req } diff --git a/e2e/agentnetwork/custom_pricing_test.go b/e2e/agentnetwork/custom_pricing_test.go new file mode 100644 index 000000000..33788a16a --- /dev/null +++ b/e2e/agentnetwork/custom_pricing_test.go @@ -0,0 +1,633 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "fmt" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// The mock vLLM upstream (harness/vllm.go) always answers with this fixed usage +// block, so every request drives deterministic token counts regardless of the +// model the client asks for. The proxy prices off the REQUEST model, not the +// upstream response model, so a made-up model id billed at operator rates lets +// these tests assert exact costs without a real vendor key. +const ( + vllmPromptTokens = 11 + vllmCompletionTokens = 2 +) + +// pricedEnv is a connected single-provider agent-network deployment pointed at +// the mock vLLM upstream, with the proxy and client up and the endpoint resolved +// — ready to drive chat. All containers are torn down via t.Cleanup. +type pricedEnv struct { + providerID string + groupID string // source group of the policy; the client peer's auto-group + policyID string // policy that authorises (and meters) the requests + upstream string // provider upstream URL, needed to re-send on a PUT update + endpoint string + proxyIP string + client *harness.Client + proxy *harness.Proxy +} + +// provisionPricedProvider brings up the full path for a cost test: a mock vLLM +// upstream, a group + reusable setup key, one openai_api provider pointed at the +// mock enumerating exactly the given models (with the operator's per-1k prices), +// a policy whose token limit switches on usage metering, and a connected proxy + +// client. The provider is created with the given models so the router dispatches +// them to this provider and the cost meter bills at these rates. +// +// Passing nil models makes it a gateway-style catch-all: the router claims every +// model, and since the synthesizer ships no per-provider-record pricing entry +// for a provider that enumerates nothing, the shipped defaults table is the only +// thing that can price the request. The policy sets no model guardrail, so the +// proxy's per-provider allowlist backstop stays empty and any model routes. +func provisionPricedProvider(t *testing.T, ctx context.Context, name string, models []api.AgentNetworkProviderModel) pricedEnv { + t.Helper() + + vllm, err := harness.StartVLLM(ctx, srv) + require.NoError(t, err, "start mock vLLM upstream") + t.Cleanup(func() { _ = vllm.Terminate(context.Background()) }) + + grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-price-" + name}) + require.NoError(t, err, "create group") + t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grp.Id) }) + + ephemeral := false + sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{ + Name: "e2e-price-" + name + "-client", + Type: "reusable", + ExpiresIn: 86400, + UsageLimit: 0, + AutoGroups: []string{grp.Id}, + Ephemeral: &ephemeral, + }) + require.NoError(t, err, "mint setup key") + require.NotEmpty(t, sk.Key, "setup key plaintext") + + // The mock ignores auth, so a dummy key satisfies the "Bearer ${API_KEY}" + // template. openai_api is a known catalog provider; the enumerated model id + // need NOT be in the catalog — the operator names it and prices it here. + dummyKey := "sk-price-e2e" + prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{ + Name: name, + ProviderId: "openai_api", + UpstreamUrl: vllm.URL, + ApiKey: &dummyKey, + Enabled: ptr(true), + BootstrapCluster: ptr(harness.AgentNetworkCluster), + Models: &models, + }) + require.NoError(t, err, "create provider") + t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) }) + + // Uncapped token limit: never blocks the handful of tokens driven here, but + // switches on usage metering — the switch that makes consumption rows record. + enabled := true + pol, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{ + Name: "e2e-price-" + name, + Enabled: &enabled, + SourceGroups: []string{grp.Id}, + DestinationProviderIds: []string{prov.Id}, + Limits: &api.AgentNetworkPolicyLimits{ + TokenLimit: api.AgentNetworkPolicyTokenLimit{ + Enabled: true, + GroupCap: 10_000_000, + UserCap: 10_000_000, + WindowSeconds: 60, + }, + }, + }) + require.NoError(t, err, "create policy") + t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), pol.Id) }) + + settings, err := srv.GetSettings(ctx) + require.NoError(t, err, "read settings") + require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned") + + proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-price-"+name+"-proxy") + require.NoError(t, err, "mint proxy token") + px, err := harness.StartProxy(ctx, srv, proxyToken) + require.NoError(t, err, "start proxy") + t.Cleanup(func() { _ = px.Terminate(context.Background()) }) + + cl, err := harness.StartClient(ctx, srv, sk.Key) + require.NoError(t, err, "start client") + t.Cleanup(func() { _ = cl.Terminate(context.Background()) }) + + require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management") + // Probe first: the GET resolves the endpoint and its first packet wakes the + // lazy proxy peer, so WaitProxyPeer then observes it connected. + proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint) + require.NoError(t, err, "resolve endpoint to proxy IP") + if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil { + t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background())) + } + + return pricedEnv{ + providerID: prov.Id, + groupID: grp.Id, + policyID: pol.Id, + upstream: vllm.URL, + endpoint: settings.Endpoint, + proxyIP: proxyIP, + client: cl, + proxy: px, + } +} + +// chatOnce drives one OpenAI-shaped chat for model through the tunnel, retrying +// to absorb first-call tunnel/DNS jitter, and returns the response body. +func chatOnce(t *testing.T, ctx context.Context, env pricedEnv, model, sessionID string) string { + t.Helper() + var code int + var body string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + c, b, cerr := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, model, "Reply with exactly: pong", sessionID) + if cerr == nil { + code, body = c, b + if code == 200 { + break + } + } + time.Sleep(5 * time.Second) + } + require.Equal(t, 200, code, + "chat for %s must return 200; body: %s\n=== proxy logs ===\n%s", model, body, env.proxy.Logs(context.Background())) + return body +} + +// findAccessLogBySession polls the access-log page for the row carrying sessionID. +func findAccessLogBySession(t *testing.T, ctx context.Context, sessionID string) api.AgentNetworkAccessLog { + t.Helper() + var row api.AgentNetworkAccessLog + require.Eventually(t, func() bool { + logs, lerr := srv.ListAccessLogs(ctx) + if lerr != nil { + return false + } + for _, r := range logs.Data { + if r.SessionId != nil && *r.SessionId == sessionID { + row = r + return true + } + } + return false + }, 30*time.Second, 2*time.Second, "session id %q must be recorded in an access-log row", sessionID) + return row +} + +// assertOpenAICostAtRates asserts an access-log row's token counts and every cost +// bucket match the mock's fixed usage priced at the given operator rates. The +// openai surface has no cache-write bucket and the mock reports no cache tokens, +// so the whole cost is input + output; cache costs must be exactly zero. +func assertOpenAICostAtRates(t *testing.T, row api.AgentNetworkAccessLog, inRate, outRate float64) { + t.Helper() + wantInput := float64(vllmPromptTokens) / 1000 * inRate + wantOutput := float64(vllmCompletionTokens) / 1000 * outRate + wantTotal := wantInput + wantOutput + + model := "" + if row.Model != nil { + model = *row.Model + } + t.Logf("[cost] model=%s in=%d out=%d rates in/out=%.4f/%.4f stored input/output/total=$%.6f/$%.6f/$%.6f expected input/output/total=$%.6f/$%.6f/$%.6f", + model, row.InputTokens, row.OutputTokens, inRate, outRate, + row.InputCostUsd, row.OutputCostUsd, row.CostUsd, wantInput, wantOutput, wantTotal) + + assert.EqualValues(t, vllmPromptTokens, row.InputTokens, "prompt tokens from the mock usage block") + assert.EqualValues(t, vllmCompletionTokens, row.OutputTokens, "completion tokens from the mock usage block") + assert.InDeltaf(t, wantInput, row.InputCostUsd, 1e-6, "input_cost_usd must be prompt tokens at the operator input rate") + assert.InDeltaf(t, wantOutput, row.OutputCostUsd, 1e-6, "output_cost_usd must be completion tokens at the operator output rate") + assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "cost_usd must be the sum of the priced buckets") + assert.Zerof(t, row.CachedInputCostUsd, "no cache-read tokens, so cached_input_cost_usd must be 0") + assert.Zerof(t, row.CacheCreationCostUsd, "openai surface has no cache-write bucket, so cache_creation_cost_usd must be 0") + assert.Zerof(t, row.CacheCostUsd, "no cache usage, so cache_cost_usd must be 0") + assert.InDeltaf(t, row.InputCostUsd+row.OutputCostUsd, row.CostUsd, 1e-9, "stored buckets must sum to cost_usd") +} + +// verifyUsageRowForSession re-checks the persisted usage row for a session +// directly in the management sqlite store — the same audit an operator runs on a +// production store.db — asserting its cost buckets match the operator rates. +func verifyUsageRowForSession(t *testing.T, sessionID string, inRate, outRate float64) { + t.Helper() + dbPath, err := srv.SnapshotStoreDB(t.TempDir()) + require.NoError(t, err, "snapshot management sqlite store") + + db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{}) + require.NoError(t, err, "open store snapshot") + sqlDB, err := db.DB() + require.NoError(t, err) + defer func() { _ = sqlDB.Close() }() + + var provider, model string + var inTok, outTok, cachedTok, cacheCreateTok int64 + var inCost, cachedInCost, cacheCreateCost, outCost float64 + row := db.Raw(`SELECT provider, model, input_tokens, output_tokens, cached_input_tokens, cache_creation_tokens, + input_cost_usd, cached_input_cost_usd, cache_creation_cost_usd, output_cost_usd + FROM agent_network_request_usage WHERE session_id = ? ORDER BY timestamp DESC LIMIT 1`, sessionID).Row() + require.NoError(t, row.Scan(&provider, &model, &inTok, &outTok, &cachedTok, &cacheCreateTok, + &inCost, &cachedInCost, &cacheCreateCost, &outCost), + "a usage row must exist for session %q", sessionID) + + wantInput := float64(inTok) / 1000 * inRate + wantOutput := float64(outTok) / 1000 * outRate + t.Logf("[sql] session=%s %s/%s in=%d out=%d stored input/cached/create/output=$%.6f/$%.6f/$%.6f/$%.6f", + sessionID, provider, model, inTok, outTok, inCost, cachedInCost, cacheCreateCost, outCost) + assert.EqualValues(t, vllmPromptTokens, inTok, "usage row prompt tokens") + assert.EqualValues(t, vllmCompletionTokens, outTok, "usage row completion tokens") + assert.InDeltaf(t, wantInput, inCost, 1e-6, "usage input_cost_usd must be prompt tokens at the operator input rate") + assert.InDeltaf(t, wantOutput, outCost, 1e-6, "usage output_cost_usd must be completion tokens at the operator output rate") + assert.Zerof(t, cachedInCost, "usage cached_input_cost_usd must be 0 (no cache usage)") + assert.Zerof(t, cacheCreateCost, "usage cache_creation_cost_usd must be 0 (no cache usage)") +} + +// TestCustomModelPricing proves an operator can serve a model that is NOT in +// NetBird's compiled catalog, at prices they type themselves, and that those +// operator prices drive the recorded cost end to end — access log AND usage +// ledger. The provider enumerates one made-up model id at deliberately odd rates +// (no default entry could supply them), the client requests it, and every cost +// bucket must equal the mock's fixed token counts multiplied by those rates. +func TestCustomModelPricing(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + customModel = "e2e-custom-model" // absent from the compiled catalog + inRate = 0.037 // odd rates so a stray default can't match + outRate = 0.089 + ) + + env := provisionPricedProvider(t, ctx, "custommodel", []api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRate, OutputPer1k: outRate}, + }) + + sessionID := "e2e-session-custommodel" + body := chatOnce(t, ctx, env, customModel, sessionID) + require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body) + + row := findAccessLogBySession(t, ctx, sessionID) + require.NotNil(t, row.Model, "access-log row must carry the requested model") + assert.Equal(t, customModel, *row.Model, "the row must be stamped with the requested (custom) model, not the mock's response model") + assertOpenAICostAtRates(t, row, inRate, outRate) + + // Metering: the uncapped token limit switches on usage recording, so the + // request must surface as a consumption row with positive tokens and cost. + require.Eventually(t, func() bool { + rows, lerr := srv.ListConsumption(ctx) + if lerr != nil { + return false + } + for _, r := range rows { + if r.TokensInput > 0 && r.TokensOutput > 0 && r.CostUsd > 0 { + return true + } + } + return false + }, 60*time.Second, 3*time.Second, "custom-model usage must be metered into a consumption row with positive cost") + + // Final raw-SQL audit: bypass the API and re-verify the persisted usage row. + verifyUsageRowForSession(t, sessionID, inRate, outRate) +} + +// TestPriceChangeUpdatesRecordedCost proves that changing a provider's model +// price is reflected in the cost recorded for subsequent requests — in both the +// access log and the usage ledger — while requests already priced at the old +// rate keep their original cost. The update propagates to the connected proxy +// live (a mapping push rebuilds the cost_meter chain with the new table), so no +// reconnect or restart is needed; the test polls a fresh request until the new +// rate lands. +func TestPriceChangeUpdatesRecordedCost(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + const ( + customModel = "e2e-repriced-model" + inRateA = 0.010 + outRateA = 0.020 + inRateB = 0.050 // 5x / 4x the original, so a repriced row is unmistakable + outRateB = 0.080 + ) + + env := provisionPricedProvider(t, ctx, "reprice", []api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRateA, OutputPer1k: outRateA}, + }) + + // Phase 1 — request priced at the original rate A. + sessionA := "e2e-session-reprice-a" + chatOnce(t, ctx, env, customModel, sessionA) + rowA := findAccessLogBySession(t, ctx, sessionA) + assertOpenAICostAtRates(t, rowA, inRateA, outRateA) + verifyUsageRowForSession(t, sessionA, inRateA, outRateA) + + // Change the model's price. The API key is omitted so the stored one is kept; + // the models array is re-sent with the new rates (PUT replaces the list). + // This reconciles synchronously and pushes a fresh cost_meter table to the + // already-connected proxy — no reconnect. + _, err := srv.UpdateProvider(ctx, env.providerID, api.AgentNetworkProviderRequest{ + Name: "reprice", + ProviderId: "openai_api", + UpstreamUrl: env.upstream, + Enabled: ptr(true), + Models: &[]api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRateB, OutputPer1k: outRateB}, + }, + }) + require.NoError(t, err, "update provider price") + + // Phase 2 — the push + chain rebuild is async, so drive fresh requests (each + // under its own session) until one is priced at the new rate B. Each iteration + // fires one request and waits for that session's row to be ingested before + // reading its cost, so an un-ingested row is never mistaken for "still rate A". + // The expected new input cost is unmistakably higher than rate A, so a + // lingering old-rate row can't satisfy the check. + wantInputB := float64(vllmPromptTokens) / 1000 * inRateB + var repriced api.AgentNetworkAccessLog + var lastSession string + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + lastSession = fmt.Sprintf("e2e-session-reprice-b-%d", time.Now().UnixNano()) + code, _, cerr := env.client.Chat(ctx, env.endpoint, env.proxyIP, harness.WireChat, customModel, "Reply with exactly: pong", lastSession) + if cerr != nil || code != 200 { + time.Sleep(5 * time.Second) + continue + } + row := findAccessLogBySession(t, ctx, lastSession) + if inDelta(row.InputCostUsd, wantInputB, 1e-6) { + repriced = row + break + } + // Still priced at the old rate — the push hasn't landed yet; retry. + time.Sleep(5 * time.Second) + } + require.NotEmpty(t, repriced.Id, "a request after the price change must be priced at the new rate B; last input_cost_usd=$%.6f, wanted $%.6f\n=== proxy logs ===\n%s", + repriced.InputCostUsd, wantInputB, env.proxy.Logs(context.Background())) + + assertOpenAICostAtRates(t, repriced, inRateB, outRateB) + verifyUsageRowForSession(t, lastSession, inRateB, outRateB) + + // The original request keeps its original cost: repricing is not retroactive. + rowAStill := findAccessLogBySession(t, ctx, sessionA) + assertOpenAICostAtRates(t, rowAStill, inRateA, outRateA) + verifyUsageRowForSession(t, sessionA, inRateA, outRateA) +} + +// TestPricingDefaultsFileDrivesCost proves the operator-supplied pricing +// defaults file is what the proxy bills with. The harness configures +// server.agentNetwork.pricingDefaultsFile as a BARE FILENAME and writes that +// file into the bind-mounted datadir (see harness.PricingDefaultsFileName), so a +// pass exercises the whole chain: combined yaml → ToManagementConfig → +// pricing.LoadFile (relative path resolved against datadir) → DefaultTable → +// the synthesizer's cost_meter defaults payload → the proxy's lookup. +// +// The provider enumerates NO models, so it is a catch-all route with no +// per-provider-record pricing entry at all — the only rates that can price the +// request are the shipped defaults. The model is a real catalog model whose +// built-in rates the file replaces with deliberately odd values, so billing at +// the compiled-in rates (i.e. the file never loaded) fails the assertions. +func TestPricingDefaultsFileDrivesCost(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + // nil models: a gateway-style provider claiming every model. The synthesizer + // ships no per-record entry for it, so the defaults table is its price list. + env := provisionPricedProvider(t, ctx, "defaultsfile", nil) + + sessionID := "e2e-session-defaultsfile" + body := chatOnce(t, ctx, env, harness.PricedDefaultModel, sessionID) + require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body) + + row := findAccessLogBySession(t, ctx, sessionID) + require.NotNil(t, row.Model, "access-log row must carry the requested model") + assert.Equal(t, harness.PricedDefaultModel, *row.Model, "the row must be stamped with the requested model") + + // The file's rates, not the compiled-in catalog rates for this model. + assertOpenAICostAtRates(t, row, harness.PricedDefaultInputPer1k, harness.PricedDefaultOutputPer1k) + verifyUsageRowForSession(t, sessionID, harness.PricedDefaultInputPer1k, harness.PricedDefaultOutputPer1k) +} + +// TestPricingDefaultsFileLeavesOtherModelsAlone proves the defaults file merges +// per entry rather than replacing the whole table: the file names exactly one +// model, so a DIFFERENT catalog model must still bill at its compiled-in rates. +// Without this, a file that shipped as a wholesale replacement would silently +// zero-cost every model the operator didn't list, and TestPricingDefaultsFile- +// DrivesCost alone would not notice. +func TestPricingDefaultsFileLeavesOtherModelsAlone(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + // gpt-4o-mini is a catalog model the pricing file does NOT mention, so it must + // keep its built-in rates. Pinned here independently of the catalog source so + // a rate change in either place surfaces as a failure to reconcile rather + // than passing silently. + const ( + untouchedModel = "gpt-4o-mini" + builtinInRate = 0.00015 + builtinOutRate = 0.0006 + ) + + env := provisionPricedProvider(t, ctx, "defaultsfileother", nil) + + sessionID := "e2e-session-defaultsfile-other" + chatOnce(t, ctx, env, untouchedModel, sessionID) + + row := findAccessLogBySession(t, ctx, sessionID) + assertOpenAICostAtRates(t, row, builtinInRate, builtinOutRate) + verifyUsageRowForSession(t, sessionID, builtinInRate, builtinOutRate) +} + +// TestCustomModelAccessLogAttribution proves a custom (non-catalog) model is +// handled correctly in the ACCESS LOG, not just in the cost columns. The other +// tests here assert money; this one asserts the row's identity and attribution +// dimensions — the columns the dashboard filters, groups and drills down on. +// +// A custom model id is the interesting case precisely because nothing in +// NetBird's catalog describes it. Its provider vendor, parser surface, cost +// buckets, and dashboard filterability all have to come from the operator's +// provider record rather than from a compiled-in entry. So this checks: +// +// - the row is stamped with the REQUESTED model id verbatim, not the mock +// upstream's response model (Qwen/Qwen2.5-0.5B-Instruct) and not a +// normalized or catalog-substituted id; +// - provider is the vendor SURFACE ("openai", from the catalog entry's +// ParserID) — a custom model does not change which wire shape was spoken; +// - resolved_provider_id / selected_policy_id / group_ids attribute the row to +// the operator's provider record, the authorising policy, and the caller's +// group, so spend on a custom model is attributable; +// - decision is "allow" with no deny reason, and the request dimensions +// (status 200, POST, the OpenAI chat path, non-stream, source IP, duration) +// are recorded; +// - management's SERVER-SIDE model filter finds the row by its custom id, so +// the model column is genuinely indexed and queryable rather than merely +// stored; +// - prompt/completion capture stays empty, since prompt collection is off by +// default and a custom model must not bypass that gate. +func TestCustomModelAccessLogAttribution(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute) + defer cancel() + + // A model id no catalog entry carries, at odd rates so its cost cannot come + // from anywhere but the provider record. + const ( + customModel = "e2e-attribution-model-v9" + inRate = 0.0271 + outRate = 0.0913 + ) + + env := provisionPricedProvider(t, ctx, "attribution", []api.AgentNetworkProviderModel{ + {Id: customModel, InputPer1k: inRate, OutputPer1k: outRate}, + }) + + sessionID := "e2e-session-attribution" + body := chatOnce(t, ctx, env, customModel, sessionID) + require.Contains(t, body, "chat.completion", "body should be an OpenAI-compatible completion; got: %s", body) + + row := findAccessLogBySession(t, ctx, sessionID) + + // Identity: the requested model verbatim. The mock answers with its own + // served model id, so a row carrying that instead means the log is sourced + // from the response body rather than the parsed request. + require.NotNil(t, row.Model, "access-log row must carry the requested model") + assert.Equal(t, customModel, *row.Model, + "the row must be stamped with the requested custom model id verbatim, not the mock upstream's response model (%s)", harness.VLLMModel) + + // Surface: a custom model id does not change the wire shape that was spoken. + // provider is the vendor surface from the catalog entry's parser, which is + // also the key the cost meter's cache formula switches on. + require.NotNil(t, row.Provider, "access-log row must carry the vendor surface") + assert.Equal(t, "openai", *row.Provider, + "openai_api's parser surface is openai, regardless of how exotic the model id is") + + // Attribution: which provider record served it, which policy authorised it, + // and which group the authorisation came through. Without these, spend on a + // custom model can be seen but not attributed. + require.NotNil(t, row.ResolvedProviderId, "row must name the provider record that served the request") + assert.Equal(t, env.providerID, *row.ResolvedProviderId, + "the router stamps the operator's provider record id; a custom model must attribute to the record that enumerated it") + require.NotNil(t, row.SelectedPolicyId, "row must name the policy that authorised the request") + assert.Equal(t, env.policyID, *row.SelectedPolicyId, + "the policy carrying the token limit is the one that paid for the request") + require.NotNil(t, row.GroupIds, "row must carry the authorising group ids") + assert.Contains(t, *row.GroupIds, env.groupID, + "the caller's group is the policy's source group, so it must be the authorising group") + + // Decision + request dimensions. + require.NotNil(t, row.Decision, "row must carry the policy decision") + assert.Equal(t, "allow", *row.Decision, "the uncapped policy allows this request") + if row.DenyReason != nil { + assert.Empty(t, *row.DenyReason, "an allowed request must carry no deny reason") + } + assert.Equal(t, 200, row.StatusCode, "the mock upstream answers 200") + if row.Method != nil { + assert.Equal(t, "POST", *row.Method, "a chat completion is a POST") + } + require.NotNil(t, row.Path, "row must record the request path") + assert.Equal(t, "/v1/chat/completions", *row.Path, + "the OpenAI chat path the client called, as seen by the proxy") + require.NotNil(t, row.Host, "row must record the host the client addressed") + assert.Equal(t, env.endpoint, *row.Host, "the agent-network endpoint the client resolved") + if row.Stream != nil { + assert.False(t, *row.Stream, "the harness sends a non-streaming request") + } + require.NotNil(t, row.SourceIp, "row must record the caller's tunnel IP") + assert.NotEmpty(t, *row.SourceIp, "the request arrived over the tunnel, so a source IP is known") + + // Tokens and cost, so the attribution above is anchored to a real priced row + // rather than an empty shell that happens to carry the right ids. + assertOpenAICostAtRates(t, row, inRate, outRate) + assert.EqualValues(t, vllmPromptTokens+vllmCompletionTokens, row.TotalTokens, + "total_tokens is the mock's reported total") + + // Prompt capture is off by default (account master switch), and a custom + // model must not bypass that gate. + if row.RequestPrompt != nil { + assert.Empty(t, *row.RequestPrompt, "prompt collection is off by default, so no prompt may be stored") + } + if row.ResponseCompletion != nil { + assert.Empty(t, *row.ResponseCompletion, "prompt collection is off by default, so no completion may be stored") + } + + // Queryability: management's SERVER-SIDE model filter must find the row by + // its custom id. findAccessLogBySession above scans a page client-side, so + // this is the check that the model column is actually indexed and filterable + // — the dashboard's per-model drill-down on a custom model depends on it. + filtered, err := srv.ListAccessLogsFiltered(ctx, url.Values{"model": []string{customModel}}) + require.NoError(t, err, "filter access logs by the custom model id") + require.Positive(t, filtered.TotalRecords, "the custom model must be findable via the server-side model filter") + foundSession := false + for _, r := range filtered.Data { + require.NotNil(t, r.Model, "filtered row must carry a model") + assert.Equal(t, customModel, *r.Model, "the model filter must not return rows for other models") + if r.SessionId != nil && *r.SessionId == sessionID { + foundSession = true + } + } + assert.True(t, foundSession, "the filtered page must include this test's request") + + // Final raw-SQL audit of the parallel usage row: the ledger must carry the + // same custom model, surface, and provider-record attribution as the log. + verifyUsageAttributionForSession(t, sessionID, customModel, "openai", env.providerID, env.groupID) +} + +// verifyUsageAttributionForSession checks the usage ledger's attribution columns +// for a session directly in the management sqlite store — including the group +// child row, which the API renders but which only exists if the proxy's +// authorising-group CSV was parsed into normalised rows. The usage table is +// written unconditionally (independent of the log-collection toggle), so this is +// the record that must attribute spend even for accounts with logs off. +func verifyUsageAttributionForSession(t *testing.T, sessionID, wantModel, wantProvider, wantProviderID, wantGroupID string) { + t.Helper() + dbPath, err := srv.SnapshotStoreDB(t.TempDir()) + require.NoError(t, err, "snapshot management sqlite store") + + db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{}) + require.NoError(t, err, "open store snapshot") + sqlDB, err := db.DB() + require.NoError(t, err) + defer func() { _ = sqlDB.Close() }() + + var id, provider, model, resolvedProviderID, userID string + require.NoError(t, db.Raw( + `SELECT id, provider, model, resolved_provider_id, user_id + FROM agent_network_request_usage WHERE session_id = ? ORDER BY timestamp DESC LIMIT 1`, sessionID). + Row().Scan(&id, &provider, &model, &resolvedProviderID, &userID), + "a usage row must exist for session %q", sessionID) + + t.Logf("[sql] usage attribution session=%s id=%s provider=%s model=%s resolved_provider_id=%s user_id=%s", + sessionID, id, provider, model, resolvedProviderID, userID) + assert.Equal(t, wantModel, model, "usage row must carry the requested custom model") + assert.Equal(t, wantProvider, provider, "usage row must carry the vendor surface") + assert.Equal(t, wantProviderID, resolvedProviderID, "usage row must attribute to the operator's provider record") + assert.NotEmpty(t, userID, "the tunnel peer resolves to a principal, so the usage row must be attributable to it") + + // The authorising group lands in the normalised child table, which is what + // the usage overview joins on to break spend down by group. + var groupIDs []string + require.NoError(t, db.Raw( + `SELECT group_id FROM agent_network_request_usage_group WHERE usage_id = ?`, id). + Scan(&groupIDs).Error, "read usage group child rows") + assert.Contains(t, groupIDs, wantGroupID, + "the authorising group must be normalised into a usage_group row so spend can be grouped by it") +} + +// inDelta reports whether a and b are within tol of each other. +func inDelta(a, b, tol float64) bool { + d := a - b + if d < 0 { + d = -d + } + return d <= tol +} diff --git a/e2e/harness/agentnetwork.go b/e2e/harness/agentnetwork.go index 53aa8e342..078e697af 100644 --- a/e2e/harness/agentnetwork.go +++ b/e2e/harness/agentnetwork.go @@ -9,6 +9,7 @@ import ( "fmt" "io" "net/http" + "net/url" "github.com/netbirdio/netbird/shared/management/http/api" ) @@ -74,6 +75,13 @@ func (c *Combined) DeleteProvider(ctx context.Context, id string) error { return anDelete(ctx, c, "/api/agent-network/providers/"+id) } +// UpdateProvider replaces a provider by id (PUT). The API key may be omitted on +// the request to keep the stored one; Models replaces the enumerated list, so +// this is the path a test uses to change a model's price mid-run. +func (c *Combined) UpdateProvider(ctx context.Context, id string, req api.AgentNetworkProviderRequest) (api.AgentNetworkProvider, error) { + return anRequest[api.AgentNetworkProvider](ctx, c, http.MethodPut, "/api/agent-network/providers/"+id, req) +} + // SetProviderEnabled toggles a provider's enabled flag, preserving its other // fields (the API key is omitted, which keeps the stored one). Used to run one // provider at a time so model→provider routing is unambiguous. @@ -139,3 +147,16 @@ func (c *Combined) ListConsumption(ctx context.Context) ([]api.AgentNetworkConsu func (c *Combined) ListAccessLogs(ctx context.Context) (api.AgentNetworkAccessLogsResponse, error) { return anRequest[api.AgentNetworkAccessLogsResponse](ctx, c, http.MethodGet, "/api/agent-network/access-logs", nil) } + +// ListAccessLogsFiltered returns the access-log page narrowed by the given +// query parameters (e.g. model=..., session_id=..., provider_id=...). This +// exercises management's server-side filtering rather than filtering client +// side, so a row that is ingested but not indexed under the filtered column +// surfaces as an empty page. +func (c *Combined) ListAccessLogsFiltered(ctx context.Context, query url.Values) (api.AgentNetworkAccessLogsResponse, error) { + path := "/api/agent-network/access-logs" + if encoded := query.Encode(); encoded != "" { + path += "?" + encoded + } + return anRequest[api.AgentNetworkAccessLogsResponse](ctx, c, http.MethodGet, path, nil) +} diff --git a/e2e/harness/combined.go b/e2e/harness/combined.go index 5723100ca..b2f0d89d2 100644 --- a/e2e/harness/combined.go +++ b/e2e/harness/combined.go @@ -93,10 +93,19 @@ func StartCombined(ctx context.Context) (*Combined, error) { _ = net.Remove(ctx) return nil, fmt.Errorf("write combined config: %w", err) } - if err := os.MkdirAll(filepath.Join(workDir, "data"), 0o755); err != nil { + dataDir := filepath.Join(workDir, "data") + if err := os.MkdirAll(dataDir, 0o755); err != nil { _ = net.Remove(ctx) return nil, fmt.Errorf("create datadir: %w", err) } + // The config's agentNetwork.pricingDefaultsFile is a bare filename, so the + // server resolves it against the datadir; write it there. It is an explicitly + // configured path, so a failure to load fails the server's startup — which + // surfaces here as the /api/instance readiness wait timing out. + if err := os.WriteFile(filepath.Join(dataDir, PricingDefaultsFileName), []byte(pricingDefaultsYAML), 0o644); err != nil { //nolint:gosec // non-secret config, bind-mounted and read by the container + _ = net.Remove(ctx) + return nil, fmt.Errorf("write pricing defaults: %w", err) + } req := testcontainers.ContainerRequest{ Image: combinedImage, diff --git a/e2e/harness/config.go b/e2e/harness/config.go index b4bed60a2..71b3656c5 100644 --- a/e2e/harness/config.go +++ b/e2e/harness/config.go @@ -8,6 +8,13 @@ package harness // embedded IdP, local signal/relay/STUN, and a sqlite store under the mounted // data dir. exposedAddress is the address peers use to reach this container; it // is overridden per-run so the value matches the container's network alias. +// +// pricingDefaultsFile is deliberately a BARE FILENAME, not an absolute path: it +// must resolve against dataDir (→ /nb/data/), which is the resolution rule +// the combined server applies. It is also an EXPLICITLY configured path, so the +// server is required to load it — a broken path or malformed file fails startup +// rather than silently falling back to the compiled-in rates, and TestMain then +// fails with the container logs. const combinedConfigYAML = `server: listenAddress: ":8080" exposedAddress: "%s" @@ -23,4 +30,36 @@ const combinedConfigYAML = `server: issuer: "%s" store: engine: "sqlite" + agentNetwork: + pricingDefaultsFile: "` + PricingDefaultsFileName + `" +` + +const ( + // PricingDefaultsFileName is the basename of the operator-supplied LLM + // pricing defaults file the combined server is configured to load. Written + // into the bind-mounted datadir by StartCombined. + PricingDefaultsFileName = "e2e_llm_pricing.yaml" + + // PricedDefaultModel is a real catalog model (openai surface) whose rates the + // defaults file below REPLACES. Tests drive it against the mock vLLM upstream + // and assert the file's rates were billed, which is only true if the file + // travelled: config → LoadFile → DefaultTable → synthesizer → the proxy's + // cost_meter defaults table. + PricedDefaultModel = "gpt-4.1-mini" + // PricedDefaultInputPer1k / PricedDefaultOutputPer1k are deliberately odd + // values that no compiled-in catalog entry carries (gpt-4.1-mini ships as + // 0.0004 / 0.0016), so a test asserting them cannot pass on the built-in + // table. + PricedDefaultInputPer1k = 0.0123 + PricedDefaultOutputPer1k = 0.0456 +) + +// pricingDefaultsYAML is the operator-supplied pricing defaults file. Its schema +// is surface -> model -> per-1k rates. Entries replace the compiled-in entry for +// the same surface+model whole; every other model keeps its built-in rates, so +// this file overriding one model must not disturb the rest of the table. +const pricingDefaultsYAML = `openai: + gpt-4.1-mini: + input_per_1k: 0.0123 + output_per_1k: 0.0456 ` diff --git a/management/cmd/management.go b/management/cmd/management.go index 79c838ec4..147985314 100644 --- a/management/cmd/management.go +++ b/management/cmd/management.go @@ -23,6 +23,7 @@ import ( "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/formatter/hook" + agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" "github.com/netbirdio/netbird/management/internals/server" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" nbdomain "github.com/netbirdio/netbird/shared/management/domain" @@ -112,6 +113,29 @@ var ( mgmtSingleAccModeDomain = "" } + // Load the management-side LLM pricing defaults file: an + // explicitly configured path is required to load (a typo must + // fail startup — the operator believes those rates are live); + // otherwise /defaults_llm_pricing.yaml is probed and + // may be absent (compiled-in defaults serve). A relative path + // is resolved against the datadir so a bare filename lands + // alongside the store. Either way the path stays watched: the + // reloader picks up edits — and the file appearing later — + // without a restart. + pricingPath := config.AgentNetwork.PricingDefaultsFile + pricingRequired := pricingPath != "" + if !pricingRequired { + pricingPath = agentnetworkpricing.DefaultFileName + } + if !filepath.IsAbs(pricingPath) { + pricingPath = filepath.Join(config.Datadir, pricingPath) + } + log.Infof("loading agent-network pricing defaults from %s (required: %v)", pricingPath, pricingRequired) + if err := agentnetworkpricing.LoadFile(pricingPath, pricingRequired); err != nil { + return fmt.Errorf("load agent-network pricing defaults: %v", err) + } + agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval) + srv := newServer(&server.Config{ NbConfig: config, DNSDomain: dnsDomain, diff --git a/management/internals/modules/agentnetwork/catalog/catalog.go b/management/internals/modules/agentnetwork/catalog/catalog.go index f82e94bf3..2c4efd0b4 100644 --- a/management/internals/modules/agentnetwork/catalog/catalog.go +++ b/management/internals/modules/agentnetwork/catalog/catalog.go @@ -7,12 +7,28 @@ package catalog import "github.com/netbirdio/netbird/shared/management/http/api" // Model is the in-memory representation of a catalog model. +// +// The three cache rates mirror the proxy cost meter's Entry semantics +// (USD per 1k tokens; 0 = no rate configured, that bucket bills at +// InputPer1k): +// - CachedInputPer1k: OpenAI-shape rate for cached prompt tokens +// (a SUBSET of input tokens). Typically 0.1-0.5x input. +// - CacheReadPer1k / CacheCreationPer1k: Anthropic-shape rates for +// the two ADDITIVE prompt-cache buckets. Typically 0.1x / 1.25x +// input. +// +// The catalog is the single default-pricing source: the agentnetwork +// pricing package folds these models into per-surface tables that the +// synthesizer ships to the proxy's cost_meter. type Model struct { - ID string - Label string - InputPer1k float64 - OutputPer1k float64 - ContextWindow int + ID string + Label string + InputPer1k float64 + OutputPer1k float64 + CachedInputPer1k float64 + CacheReadPer1k float64 + CacheCreationPer1k float64 + ContextWindow int } // ProviderKind groups catalog entries for UI presentation. The split @@ -65,6 +81,17 @@ type Provider struct { // surface — the proxy middleware then falls back to URL sniffing // or skips request-side enrichment. ParserID string + // PricingSurfaces names the cost-meter pricing surfaces this + // provider's Models are priced under ("openai", "anthropic", + // "bedrock" — the llm.Parser surface the request parser stamps as + // llm.provider at billing time). NOT derivable from ParserID: + // bedrock_api and vertex_ai_api leave ParserID empty (URL-sniffed) + // yet price under "bedrock" / "anthropic", and kimi_api serves two + // body shapes so it prices under both. Nil for gateway/custom + // entries, which declare no models. Same (surface, model) pair + // contributed by two providers must carry identical rates — the + // pricing package's tests enforce that. + PricingSurfaces []string // IdentityInjection, when non-nil, instructs the proxy to stamp // the caller's NetBird identity onto upstream requests under the // configured header names. Used for gateways like LiteLLM that @@ -219,6 +246,7 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#10A37F", ParserID: "openai", + PricingSurfaces: []string{"openai"}, // Pricing + context windows cross-checked against LiteLLM's // model_prices_and_context_window.json. Notable corrections from // earlier values: o4-mini repriced from $4/$16 to $1.10/$4.40 @@ -226,20 +254,20 @@ var providers = []Provider{ // family context windows split between 1.05M for full-size // models and 272K for mini/nano/codex variants. Models: []Model{ - {ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000}, - {ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000}, - {ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000}, - {ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000}, - {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000}, - {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000}, - {ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 272000}, - {ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 128000}, - {ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000}, - {ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576}, - {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576}, - {ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, ContextWindow: 1047576}, - {ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000}, - {ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000}, + {ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000}, + {ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000}, + {ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000}, + {ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000}, + {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000}, + {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000}, + {ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 272000}, + {ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 128000}, + {ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000}, + {ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576}, + {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576}, + {ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, CachedInputPer1k: 0.000025, ContextWindow: 1047576}, + {ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000}, + {ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000}, {ID: "gpt-4-turbo", Label: "GPT-4 Turbo", InputPer1k: 0.01, OutputPer1k: 0.03, ContextWindow: 128000}, {ID: "gpt-3.5-turbo", Label: "GPT-3.5 Turbo", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385}, {ID: "text-embedding-3-large", Label: "text-embedding-3-large", InputPer1k: 0.00013, OutputPer1k: 0, ContextWindow: 8191}, @@ -257,6 +285,7 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#D97757", ParserID: "anthropic", + PricingSurfaces: []string{"anthropic"}, // Per Anthropic's current model lineup. Pricing in USD per 1k // tokens. Context windows: 4.6+ family is 1M; Haiku 4.5 stays at // 200K. claude-3-7-sonnet and claude-3-5-haiku retired @@ -267,14 +296,14 @@ var providers = []Provider{ // account to be on >= 30-day data retention or all requests // 400. Models: []Model{ - {ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000}, - {ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000}, - {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, - {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000}, - {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000}, + {ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000}, + {ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000}, + {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000}, + {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000}, + {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000}, }, }, { @@ -288,18 +317,19 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#0078D4", ParserID: "openai", + PricingSurfaces: []string{"openai"}, // Mirrors openai_api pricing — Azure resells OpenAI models at the // same per-token rates, just under different deployment names. Models: []Model{ - {ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000}, - {ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000}, - {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000}, - {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000}, - {ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000}, - {ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576}, - {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576}, - {ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000}, - {ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000}, + {ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000}, + {ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000}, + {ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000}, + {ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000}, + {ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000}, + {ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576}, + {ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576}, + {ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000}, + {ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000}, {ID: "gpt-35-turbo", Label: "GPT-3.5 Turbo (Azure)", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385}, }, }, @@ -313,6 +343,9 @@ var providers = []Provider{ AuthHeaderTemplate: "Bearer ${API_KEY}", DefaultContentType: "application/json", BrandColor: "#FF9900", + // ParserID stays empty (path-style dispatch via IsBedrockPathStyle); + // the request parser meters these under the "bedrock" surface. + PricingSurfaces: []string{"bedrock"}, // Anthropic models on Bedrock take the anthropic.* prefix and // follow the same lineup / pricing as the first-party Anthropic // catalog entry above. claude-3-7-sonnet and claude-3-5-haiku @@ -322,13 +355,13 @@ var providers = []Provider{ // Llama 3.3 70B entry kept unchanged — LiteLLM tracks only // per-region Llama 3 entries; standalone 3.3 not yet listed. Models: []Model{ - {ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000}, - {ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, - {ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000}, - {ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000}, + {ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000}, + {ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000}, + {ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000}, + {ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000}, {ID: "meta.llama3-3-70b-instruct", Label: "Llama 3.3 70B (Bedrock)", InputPer1k: 0.00072, OutputPer1k: 0.00072, ContextWindow: 128000}, {ID: "amazon.nova-2-lite", Label: "Amazon Nova 2 Lite (Bedrock, preview)", InputPer1k: 0.0003, OutputPer1k: 0.0025, ContextWindow: 1000000}, {ID: "amazon.nova-pro", Label: "Amazon Nova Pro (Bedrock)", InputPer1k: 0.0008, OutputPer1k: 0.0032, ContextWindow: 300000}, @@ -358,6 +391,10 @@ var providers = []Provider{ AuthHeaderTemplate: "Bearer ${API_KEY}", DefaultContentType: "application/json", BrandColor: "#4285F4", + // ParserID stays empty (path-style dispatch via IsVertexPathStyle); + // Anthropic-on-Vertex requests are metered under the "anthropic" + // surface with the bare, unversioned model id. + PricingSurfaces: []string{"anthropic"}, // Vertex carries the model in the URL path and authenticates with a // service-account-minted OAuth token (api_key = "keyfile::"). // Only Anthropic-on-Vertex is metered today: the request parser maps the @@ -369,14 +406,14 @@ var providers = []Provider{ // exists — the router denies unmeterable publishers rather than forward // them uncounted. Models: []Model{ - {ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000}, - {ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000}, - {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000}, - {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, - {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000}, - {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000}, + {ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000}, + {ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000}, + {ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000}, + {ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000}, + {ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000}, + {ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000}, }, }, { @@ -390,6 +427,7 @@ var providers = []Provider{ DefaultContentType: "application/json", BrandColor: "#FF7000", ParserID: "openai", + PricingSurfaces: []string{"openai"}, // Pricing + context windows cross-checked against LiteLLM. Key // gotchas the marketing page hides: // - `mistral-medium-latest` aliases to Medium 3.1 ($0.40/$2), @@ -448,6 +486,10 @@ var providers = []Provider{ // model id "k3") is account-bound seat licensing rather than a // meterable platform key, so it's deliberately not the default. ParserID: "", + // Both body shapes are metered: /v1/chat/completions under + // "openai", /anthropic/v1/messages under "anthropic" — so the + // K3 entry is priced on both surfaces. + PricingSurfaces: []string{"openai", "anthropic"}, // Pricing per Moonshot's platform rates at K3 launch (July 2026): // $3/$15 per MTok with $0.30 cached input, flat across the 1M-token // window. kimi-k3 is the ONLY model the platform serves newer @@ -458,7 +500,14 @@ var providers = []Provider{ // The consumer app's "K3 Swarm Max" mode is not an API SKU, so it // doesn't appear here. Models: []Model{ - {ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000}, + // Carries both cache shapes: Moonshot reports cache hits + // OpenAI-style on /v1/chat/completions (CachedInputPer1k) + // and Anthropic-style on /anthropic/v1/messages + // (CacheReadPer1k) — $0.30/MTok either way. Each surface's + // cost formula reads only its own field, so the superset + // entry prices both endpoints correctly. No cache-creation + // rate published; writes bill at the input rate. + {ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, CachedInputPer1k: 0.0003, CacheReadPer1k: 0.0003, ContextWindow: 1000000}, }, }, { @@ -758,13 +807,28 @@ func IsBedrockPathStyle(providerID string) bool { func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider { models := make([]api.AgentNetworkCatalogModel, 0, len(p.Models)) for _, m := range p.Models { - models = append(models, api.AgentNetworkCatalogModel{ + am := api.AgentNetworkCatalogModel{ Id: m.ID, Label: m.Label, InputPer1k: m.InputPer1k, OutputPer1k: m.OutputPer1k, ContextWindow: m.ContextWindow, - }) + } + // Cache rates are emitted only when configured so the dashboard + // can prefill them; 0 stays off the wire (absent = no rate). + if m.CachedInputPer1k > 0 { + v := m.CachedInputPer1k + am.CachedInputPer1k = &v + } + if m.CacheReadPer1k > 0 { + v := m.CacheReadPer1k + am.CacheReadPer1k = &v + } + if m.CacheCreationPer1k > 0 { + v := m.CacheCreationPer1k + am.CacheCreationPer1k = &v + } + models = append(models, am) } kind := api.AgentNetworkCatalogProviderKindProvider switch p.Kind { @@ -784,6 +848,10 @@ func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider { BrandColor: p.BrandColor, Models: models, } + if len(p.PricingSurfaces) > 0 { + surfaces := append([]string(nil), p.PricingSurfaces...) + resp.PricingSurfaces = &surfaces + } if len(p.ExtraHeaders) > 0 { extras := make([]api.AgentNetworkCatalogExtraHeader, 0, len(p.ExtraHeaders)) for _, h := range p.ExtraHeaders { diff --git a/management/internals/modules/agentnetwork/handlers/providers_handler.go b/management/internals/modules/agentnetwork/handlers/providers_handler.go index 13da137d5..c05363101 100644 --- a/management/internals/modules/agentnetwork/handlers/providers_handler.go +++ b/management/internals/modules/agentnetwork/handlers/providers_handler.go @@ -7,6 +7,7 @@ package handlers import ( "encoding/json" + "math" "net/http" "net/url" "strings" @@ -15,6 +16,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/shared/management/http/api" @@ -52,11 +54,45 @@ func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) { entries := catalog.All() out := make([]api.AgentNetworkCatalogProvider, 0, len(entries)) for _, e := range entries { - out = append(out, e.ToAPIResponse()) + resp := e.ToAPIResponse() + applyDefaultPricing(e, &resp) + out = append(out, resp) } util.WriteJSONObject(r.Context(), w, out) } +// applyDefaultPricing overwrites the catalog response's model rates with +// the LIVE default pricing table, which may differ from the compiled-in +// catalog rates when the operator provides a defaults_llm_pricing.yaml. +// This keeps the dashboard's model-row prefill identical to what the +// proxy will actually bill — the same table the synthesizer ships. +func applyDefaultPricing(cp catalog.Provider, resp *api.AgentNetworkCatalogProvider) { + if len(cp.PricingSurfaces) == 0 { + return + } + for i := range resp.Models { + m := &resp.Models[i] + e, ok := pricing.LookupDefault(cp.PricingSurfaces, m.Id) + if !ok { + continue + } + m.InputPer1k = e.InputPer1k + m.OutputPer1k = e.OutputPer1k + m.CachedInputPer1k = positiveRatePtr(e.CachedInputPer1k) + m.CacheReadPer1k = positiveRatePtr(e.CacheReadPer1k) + m.CacheCreationPer1k = positiveRatePtr(e.CacheCreationPer1k) + } +} + +// positiveRatePtr renders a cache rate for the API: absent (nil) when +// unset, matching the catalog response convention. +func positiveRatePtr(v float64) *float64 { + if v <= 0 { + return nil + } + return &v +} + func (h *handler) getAllProviders(w http.ResponseWriter, r *http.Request) { userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) if err != nil { @@ -213,5 +249,38 @@ func validate(req *api.AgentNetworkProviderRequest, requireAPIKey bool) error { if requireAPIKey && (req.ApiKey == nil || strings.TrimSpace(*req.ApiKey) == "") { return status.Errorf(status.InvalidArgument, "api_key is required") } + if req.Models != nil { + for i, m := range *req.Models { + if err := validateModel(i, m); err != nil { + return err + } + } + } + return nil +} + +// validateModel is the single ingress guard for operator-entered pricing: +// these rates are synthesized into the proxy's cost_meter config verbatim, +// and a negative or non-finite rate there would poison every cost the +// proxy records, so reject at the API boundary. +func validateModel(i int, m api.AgentNetworkProviderModel) error { + if strings.TrimSpace(m.Id) == "" { + return status.Errorf(status.InvalidArgument, "models[%d]: id is required", i) + } + rates := map[string]*float64{ + "input_per_1k": &m.InputPer1k, + "output_per_1k": &m.OutputPer1k, + "cached_input_per_1k": m.CachedInputPer1k, + "cache_read_per_1k": m.CacheReadPer1k, + "cache_creation_per_1k": m.CacheCreationPer1k, + } + for field, v := range rates { + if v == nil { + continue + } + if *v < 0 || math.IsNaN(*v) || math.IsInf(*v, 0) { + return status.Errorf(status.InvalidArgument, "models[%d] (%s): %s must be a finite, non-negative USD rate", i, m.Id, field) + } + } return nil } diff --git a/management/internals/modules/agentnetwork/handlers/providers_handler_test.go b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go new file mode 100644 index 000000000..649224c02 --- /dev/null +++ b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go @@ -0,0 +1,53 @@ +package handlers + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +func f(v float64) *float64 { return &v } + +// TestValidate_ModelRates guards the single ingress point for operator-entered +// pricing. These rates flow verbatim into the proxy's cost_meter config at +// synthesis time; the proxy treats a bad rate as a chain-build failure, so +// rejecting here is what keeps an account's gateway from going down. +func TestValidate_ModelRates(t *testing.T) { + base := func(models ...api.AgentNetworkProviderModel) *api.AgentNetworkProviderRequest { + key := "sk-test" + return &api.AgentNetworkProviderRequest{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + ApiKey: &key, + Models: &models, + } + } + + valid := api.AgentNetworkProviderModel{ + Id: "gpt-4o", InputPer1k: 0.0025, OutputPer1k: 0.01, + CachedInputPer1k: f(0.00125), + } + require.NoError(t, validate(base(valid), true), "finite non-negative rates must pass") + + zeroRates := api.AgentNetworkProviderModel{Id: "self-hosted-llama", InputPer1k: 0, OutputPer1k: 0} + require.NoError(t, validate(base(zeroRates), true), "explicit zero prices are allowed (free / self-hosted models)") + + cases := map[string]api.AgentNetworkProviderModel{ + "empty id": {Id: " ", InputPer1k: 0.001, OutputPer1k: 0.002}, + "negative input": {Id: "m", InputPer1k: -0.001, OutputPer1k: 0.002}, + "negative output": {Id: "m", InputPer1k: 0.001, OutputPer1k: -0.002}, + "NaN input": {Id: "m", InputPer1k: math.NaN(), OutputPer1k: 0.002}, + "Inf output": {Id: "m", InputPer1k: 0.001, OutputPer1k: math.Inf(1)}, + "negative cached": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CachedInputPer1k: f(-1)}, + "NaN cache read": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheReadPer1k: f(math.NaN())}, + "Inf cache creation": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheCreationPer1k: f(math.Inf(-1))}, + } + for name, m := range cases { + assert.Error(t, validate(base(m), true), "case %q must be rejected", name) + } +} diff --git a/management/internals/modules/agentnetwork/pricing/defaults.go b/management/internals/modules/agentnetwork/pricing/defaults.go new file mode 100644 index 000000000..c690313bc --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/defaults.go @@ -0,0 +1,156 @@ +// Package pricing builds the default LLM pricing table the synthesizer +// ships to the proxy's cost_meter middleware. The catalog is the single +// source of default rates: every catalog provider's models are folded +// into the pricing surfaces the provider declares (PricingSurfaces), +// then a small supplemental list adds priced-but-not-operator-selectable +// entries. Management is the sole pricing authority — the proxy carries +// no embedded price list and bills exclusively from the table it is +// sent. +//go:generate go run gen.go + +package pricing + +import ( + "sync" + "sync/atomic" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" +) + +// Entry is a single model's pricing in USD per 1k tokens. This struct IS +// the wire shape: the synthesizer marshals it verbatim into cost_meter's +// ConfigJSON, and the proxy unmarshals the same field names. +// +// A zero rate means "no rate configured" — the proxy bills that cache +// bucket at InputPer1k (identical semantics to the retired proxy-embedded +// table). CachedInputPer1k is the OpenAI shape (cached prompt tokens are +// a subset of input); CacheReadPer1k / CacheCreationPer1k are the +// Anthropic shape (additive buckets). +type Entry struct { + InputPer1k float64 `json:"input_per_1k"` + OutputPer1k float64 `json:"output_per_1k"` + CachedInputPer1k float64 `json:"cached_input_per_1k,omitempty"` + CacheReadPer1k float64 `json:"cache_read_per_1k,omitempty"` + CacheCreationPer1k float64 `json:"cache_creation_per_1k,omitempty"` +} + +// supplementalDefaults are (surface, model) entries that are priced but +// deliberately not operator-selectable in the catalog. Each carries a +// reason; when one of these models joins a catalog lineup, delete the +// row here — the collision test fails loudly if the rates ever disagree. +var supplementalDefaults = map[string]map[string]Entry{ + "openai": { + // GPT-5 (2025) family — kept for gateway requests using the + // unsuffixed ids; the dashboard offers only the 5.x lineup. + "gpt-5": {InputPer1k: 0.00125, OutputPer1k: 0.01, CachedInputPer1k: 0.000125}, + "gpt-5-mini": {InputPer1k: 0.00025, OutputPer1k: 0.002, CachedInputPer1k: 0.000025}, + "gpt-5-nano": {InputPer1k: 0.00005, OutputPer1k: 0.0004, CachedInputPer1k: 0.000005}, + }, + "anthropic": { + // claude-opus-5 is not yet in the catalog lineup but gateway / + // grandfathered traffic uses it; priced so it isn't skipped. + "claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625}, + // "kimi-k3[1m]" is the 1M-context alias some Claude Code guides + // configure against Moonshot's Anthropic-compatible endpoint; + // priced identically to kimi-k3 so those requests aren't skipped. + "kimi-k3[1m]": {InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003}, + }, + "bedrock": { + "anthropic.claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625}, + }, +} + +var ( + compiledOnce sync.Once + compiledTable map[string]map[string]Entry + // mergedTable holds the current live table when a pricing defaults + // file is loaded: the file merged entry-whole over the compiled-in + // base. Nil while no file is loaded (or after the file is removed), + // in which case the compiled-in table serves. Swapped atomically by + // the file loader/reloader; readers never block. + mergedTable atomic.Pointer[map[string]map[string]Entry] +) + +// DefaultTable returns the current default pricing table keyed +// surface -> model -> Entry: the management-side defaults file (see +// LoadFile / StartReloader) when one is loaded, merged over the +// compiled-in catalog table, which alone serves as the fallback when no +// file exists. The snapshot may change between calls as the file is +// re-read — consumers (the synthesizer on every reconcile, the catalog +// endpoint on every request) pick up fresh rates automatically. Callers +// must not mutate the returned maps. +func DefaultTable() map[string]map[string]Entry { + if t := mergedTable.Load(); t != nil { + return *t + } + return compiledBase() +} + +// compiledBase returns the compiled-in table (catalog + supplementals), +// built once. +func compiledBase() map[string]map[string]Entry { + compiledOnce.Do(func() { + compiledTable = buildDefaultTable() + }) + return compiledTable +} + +func buildDefaultTable() map[string]map[string]Entry { + out := make(map[string]map[string]Entry) + for _, p := range catalog.All() { + for _, surface := range p.PricingSurfaces { + inner, ok := out[surface] + if !ok { + inner = make(map[string]Entry) + out[surface] = inner + } + for _, m := range p.Models { + // First writer wins; providers contributing the same + // (surface, model) must agree on rates — enforced by + // TestDefaultTable_NoConflictingContributions. + if _, dup := inner[m.ID]; dup { + continue + } + inner[m.ID] = entryFromCatalogModel(m) + } + } + } + for surface, models := range supplementalDefaults { + inner, ok := out[surface] + if !ok { + inner = make(map[string]Entry) + out[surface] = inner + } + for id, e := range models { + if _, dup := inner[id]; dup { + continue + } + inner[id] = e + } + } + return out +} + +func entryFromCatalogModel(m catalog.Model) Entry { + return Entry{ + InputPer1k: m.InputPer1k, + OutputPer1k: m.OutputPer1k, + CachedInputPer1k: m.CachedInputPer1k, + CacheReadPer1k: m.CacheReadPer1k, + CacheCreationPer1k: m.CacheCreationPer1k, + } +} + +// LookupDefault returns the default entry for model on the first of the +// given surfaces that prices it. Used by the synthesizer to seed a +// per-provider entry with default cache rates before overlaying the +// operator's stored prices. +func LookupDefault(surfaces []string, model string) (Entry, bool) { + table := DefaultTable() + for _, s := range surfaces { + if e, ok := table[s][model]; ok { + return e, true + } + } + return Entry{}, false +} diff --git a/proxy/internal/llm/pricing/defaults_pricing.yaml b/management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml similarity index 63% rename from proxy/internal/llm/pricing/defaults_pricing.yaml rename to management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml index 988426105..bb1cb09a8 100644 --- a/proxy/internal/llm/pricing/defaults_pricing.yaml +++ b/management/internals/modules/agentnetwork/pricing/defaults_llm_pricing.example.yaml @@ -1,83 +1,176 @@ -# Embedded default pricing for llm_observability. Compiled into the proxy -# binary via go:embed in pricing.go; cost annotation works out of the box -# without any operator action. +# Default LLM pricing used by NetBird's Agent Network cost metering. +# GENERATED from the management catalog — do not edit this file in the +# repository; regenerate with: # -# Operators override entries by dropping a pricing.yaml into --plugin-data-dir -# (or whichever basename is given via params.pricing_path). The override file -# only needs entries the operator wants to change; missing entries fall -# through to these defaults. +# go generate ./management/internals/modules/agentnetwork/pricing # -# Values are USD per 1_000 tokens. Public list prices drift; ship a fresh -# binary or override individual entries via the override file as needed. +# Operators: copy this file to /defaults_llm_pricing.yaml (or +# any path configured via management.json: # -# Optional cache fields: -# cached_input_per_1k OpenAI: rate for prompt_tokens_details.cached_tokens -# (a SUBSET of prompt_tokens). Typically 0.5x input. -# Absent → cached portion bills at input_per_1k. -# cache_read_per_1k Anthropic: rate for cache_read_input_tokens -# (ADDITIVE to input_tokens). Typically 0.1x input. -# Absent → cache reads bill at input_per_1k. -# cache_creation_per_1k Anthropic: rate for cache_creation_input_tokens -# (ADDITIVE to input_tokens). Typically 1.25x input. -# Absent → cache writes bill at input_per_1k. +# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } } +# +# ) and adjust the entries you want to change. Management re-reads the +# file periodically (mtime poll, every minute): the live table feeds the +# proxies' cost metering and the dashboard's model-price prefill, so +# edits apply without a restart. Your file only needs the entries you +# want to change — but each entry REPLACES the built-in entry for that +# surface+model whole, so repeat the cache rates you want to keep. +# Unknown fields and negative or non-finite rates are rejected: at +# startup that fails boot (for an explicitly configured path); at +# runtime the previous table is kept and a warning is logged. Deleting +# the file reverts to the built-in defaults below. +# +# Top-level keys are pricing surfaces — the parser shape requests are +# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible +# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock" +# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be +# the normalized id the proxy meters (version/region suffixes stripped). +# +# Values are USD per 1_000 tokens. Optional cache fields: +# cached_input_per_1k OpenAI shape: rate for cached prompt tokens +# (a SUBSET of input tokens). Absent -> cached +# portion bills at input_per_1k. +# cache_read_per_1k Anthropic shape: rate for cache_read tokens +# (ADDITIVE to input). Absent -> input rate. +# cache_creation_per_1k Anthropic shape: rate for cache_creation +# tokens (ADDITIVE to input). Absent -> input +# rate. + +anthropic: + claude-fable-5: + input_per_1k: 0.01 + output_per_1k: 0.05 + cache_read_per_1k: 0.001 + cache_creation_per_1k: 0.0125 + claude-haiku-4-5: + input_per_1k: 0.001 + output_per_1k: 0.005 + cache_read_per_1k: 0.0001 + cache_creation_per_1k: 0.00125 + claude-opus-4-1: + input_per_1k: 0.015 + output_per_1k: 0.075 + cache_read_per_1k: 0.0015 + cache_creation_per_1k: 0.01875 + claude-opus-4-6: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-opus-4-7: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-opus-4-8: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-opus-5: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + claude-sonnet-4-5: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + claude-sonnet-4-6: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + kimi-k3: + input_per_1k: 0.003 + output_per_1k: 0.015 + cached_input_per_1k: 0.0003 + cache_read_per_1k: 0.0003 + "kimi-k3[1m]": + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + +bedrock: + amazon.nova-2-lite: + input_per_1k: 0.0003 + output_per_1k: 0.0025 + amazon.nova-lite: + input_per_1k: 0.00006 + output_per_1k: 0.00024 + amazon.nova-micro: + input_per_1k: 0.000035 + output_per_1k: 0.00014 + amazon.nova-pro: + input_per_1k: 0.0008 + output_per_1k: 0.0032 + anthropic.claude-haiku-4-5: + input_per_1k: 0.001 + output_per_1k: 0.005 + cache_read_per_1k: 0.0001 + cache_creation_per_1k: 0.00125 + anthropic.claude-opus-4-1: + input_per_1k: 0.015 + output_per_1k: 0.075 + cache_read_per_1k: 0.0015 + cache_creation_per_1k: 0.01875 + anthropic.claude-opus-4-6: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-opus-4-7: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-opus-4-8: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-opus-5: + input_per_1k: 0.005 + output_per_1k: 0.025 + cache_read_per_1k: 0.0005 + cache_creation_per_1k: 0.00625 + anthropic.claude-sonnet-4-5: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + anthropic.claude-sonnet-4-6: + input_per_1k: 0.003 + output_per_1k: 0.015 + cache_read_per_1k: 0.0003 + cache_creation_per_1k: 0.00375 + meta.llama3-3-70b-instruct: + input_per_1k: 0.00072 + output_per_1k: 0.00072 openai: - # OpenAI + OpenAI-compatible providers (openai_api, azure_openai_api, - # mistral_api, and the openai-parser gateways) all emit llm.provider="openai", - # so their models are priced here. Kept in sync with the management catalog; - # rates cross-checked against LiteLLM model_prices_and_context_window.json. - - # GPT-5.x family — cache reads 10% of input (0.1x). - gpt-5.5: - input_per_1k: 0.005 - output_per_1k: 0.03 - cached_input_per_1k: 0.0005 - gpt-5.5-pro: - input_per_1k: 0.03 - output_per_1k: 0.18 - cached_input_per_1k: 0.003 - gpt-5.4: - input_per_1k: 0.0025 - output_per_1k: 0.015 - cached_input_per_1k: 0.00025 - gpt-5.4-pro: - input_per_1k: 0.03 - output_per_1k: 0.18 - cached_input_per_1k: 0.003 - gpt-5.4-mini: - input_per_1k: 0.00075 - output_per_1k: 0.0045 - cached_input_per_1k: 0.000075 - gpt-5.4-nano: - input_per_1k: 0.0002 - output_per_1k: 0.00125 - cached_input_per_1k: 0.00002 - gpt-5.3-codex: - input_per_1k: 0.00175 - output_per_1k: 0.014 - cached_input_per_1k: 0.000175 - gpt-5.3-chat-latest: - input_per_1k: 0.00175 - output_per_1k: 0.014 - cached_input_per_1k: 0.000175 - # GPT-5 (2025) family — kept for gateway requests using the unsuffixed ids. - gpt-5: - input_per_1k: 0.00125 - output_per_1k: 0.01 - cached_input_per_1k: 0.000125 - gpt-5-mini: - input_per_1k: 0.00025 + codestral-2508: + input_per_1k: 0.0003 + output_per_1k: 0.0009 + codestral-latest: + input_per_1k: 0.001 + output_per_1k: 0.003 + devstral-medium-latest: + input_per_1k: 0.0004 output_per_1k: 0.002 - cached_input_per_1k: 0.000025 - gpt-5-nano: - input_per_1k: 0.00005 - output_per_1k: 0.0004 - cached_input_per_1k: 0.000005 - o4-mini: - input_per_1k: 0.0011 - output_per_1k: 0.0044 - cached_input_per_1k: 0.000275 - # GPT-4.1 family — cache reads 25% of input. + devstral-small-latest: + input_per_1k: 0.0001 + output_per_1k: 0.0003 + gpt-3.5-turbo: + input_per_1k: 0.0005 + output_per_1k: 0.0015 + gpt-35-turbo: + input_per_1k: 0.0005 + output_per_1k: 0.0015 + gpt-4-turbo: + input_per_1k: 0.01 + output_per_1k: 0.03 gpt-4.1: input_per_1k: 0.002 output_per_1k: 0.008 @@ -90,7 +183,6 @@ openai: input_per_1k: 0.0001 output_per_1k: 0.0004 cached_input_per_1k: 0.000025 - # GPT-4o family — cache reads 50% of input (0.5x). gpt-4o: input_per_1k: 0.0025 output_per_1k: 0.01 @@ -99,200 +191,92 @@ openai: input_per_1k: 0.00015 output_per_1k: 0.0006 cached_input_per_1k: 0.000075 - # Older GPT — no prompt caching. - gpt-4-turbo: - input_per_1k: 0.01 - output_per_1k: 0.03 - gpt-3.5-turbo: - input_per_1k: 0.0005 - output_per_1k: 0.0015 - gpt-35-turbo: # Azure deployment alias of gpt-3.5-turbo - input_per_1k: 0.0005 - output_per_1k: 0.0015 - # Embeddings — no caching, no output tokens. - text-embedding-3-large: - input_per_1k: 0.00013 - output_per_1k: 0 - text-embedding-3-small: - input_per_1k: 0.00002 - output_per_1k: 0 - - # Mistral (mistral_api) — routed via the openai parser; no prompt caching. - mistral-large-latest: - input_per_1k: 0.0005 - output_per_1k: 0.0015 - mistral-medium-latest: - input_per_1k: 0.0004 + gpt-5: + input_per_1k: 0.00125 + output_per_1k: 0.01 + cached_input_per_1k: 0.000125 + gpt-5-mini: + input_per_1k: 0.00025 output_per_1k: 0.002 - mistral-medium-3-5: - input_per_1k: 0.0015 - output_per_1k: 0.0075 - mistral-small-latest: - input_per_1k: 0.00006 - output_per_1k: 0.00018 + cached_input_per_1k: 0.000025 + gpt-5-nano: + input_per_1k: 0.00005 + output_per_1k: 0.0004 + cached_input_per_1k: 0.000005 + gpt-5.3-chat-latest: + input_per_1k: 0.00175 + output_per_1k: 0.014 + cached_input_per_1k: 0.000175 + gpt-5.3-codex: + input_per_1k: 0.00175 + output_per_1k: 0.014 + cached_input_per_1k: 0.000175 + gpt-5.4: + input_per_1k: 0.0025 + output_per_1k: 0.015 + cached_input_per_1k: 0.00025 + gpt-5.4-mini: + input_per_1k: 0.00075 + output_per_1k: 0.0045 + cached_input_per_1k: 0.000075 + gpt-5.4-nano: + input_per_1k: 0.0002 + output_per_1k: 0.00125 + cached_input_per_1k: 0.00002 + gpt-5.4-pro: + input_per_1k: 0.03 + output_per_1k: 0.18 + cached_input_per_1k: 0.003 + gpt-5.5: + input_per_1k: 0.005 + output_per_1k: 0.03 + cached_input_per_1k: 0.0005 + gpt-5.5-pro: + input_per_1k: 0.03 + output_per_1k: 0.18 + cached_input_per_1k: 0.003 + kimi-k3: + input_per_1k: 0.003 + output_per_1k: 0.015 + cached_input_per_1k: 0.0003 + cache_read_per_1k: 0.0003 magistral-medium-latest: input_per_1k: 0.002 output_per_1k: 0.005 magistral-small-latest: input_per_1k: 0.0005 output_per_1k: 0.0015 - devstral-medium-latest: - input_per_1k: 0.0004 - output_per_1k: 0.002 - devstral-small-latest: - input_per_1k: 0.0001 - output_per_1k: 0.0003 - codestral-2508: - input_per_1k: 0.0003 - output_per_1k: 0.0009 - codestral-latest: - input_per_1k: 0.001 - output_per_1k: 0.003 ministral-3-14b-2512: input_per_1k: 0.0002 output_per_1k: 0.0002 - ministral-8b-latest: - input_per_1k: 0.00015 - output_per_1k: 0.00015 ministral-3-3b-2512: input_per_1k: 0.0001 output_per_1k: 0.0001 + ministral-8b-latest: + input_per_1k: 0.00015 + output_per_1k: 0.00015 mistral-embed: input_per_1k: 0.0001 output_per_1k: 0 - - # Kimi / Moonshot AI (kimi_api) — OpenAI-compatible /v1 endpoint. Moonshot - # reports cache hits OpenAI-style when present; cached input is 10% of - # input ($0.30 vs $3.00 per MTok). kimi-k3 is the only model the platform - # serves newer accounts (K2-era ids and kimi-latest 404), matching the - # management catalog. - kimi-k3: - input_per_1k: 0.003 - output_per_1k: 0.015 - cached_input_per_1k: 0.0003 - -anthropic: - # Claude 4.x family — cache reads ≈10% of input, cache writes ≈125% of input. - # Pricing source: Anthropic's current published rates per million tokens, - # divided by 1000 for the per-1k figures stored here. - claude-fable-5: - input_per_1k: 0.010 - output_per_1k: 0.050 - cache_read_per_1k: 0.001 - cache_creation_per_1k: 0.0125 - claude-opus-5: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-8: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-7: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-6: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - claude-opus-4-1: - input_per_1k: 0.015 - output_per_1k: 0.075 - cache_read_per_1k: 0.0015 - cache_creation_per_1k: 0.01875 - claude-sonnet-4-6: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - claude-haiku-4-5: - input_per_1k: 0.001 - output_per_1k: 0.005 - cache_read_per_1k: 0.0001 - cache_creation_per_1k: 0.00125 - - # Kimi / Moonshot AI (kimi_api) via the Anthropic-compatible endpoint - # (/anthropic/v1/messages — the official Claude Code setup). Same rates - # as the OpenAI-shape entry above. "kimi-k3[1m]" is the model id some - # Claude Code guides set for the 1M-context alias; priced identically so - # cost metering doesn't silently skip those requests. - kimi-k3: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - "kimi-k3[1m]": - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - -bedrock: - # AWS Bedrock model ids, normalised by the request parser (cross-region - # inference-profile prefix + version/throughput suffix stripped), e.g. - # eu.anthropic.claude-sonnet-4-5-20250929-v1:0 -> anthropic.claude-sonnet-4-5. - # Anthropic-on-Bedrock keeps the additive cache buckets (read ≈0.1x input, - # write ≈1.25x input); Nova / Llama report no cache, so cost is input+output. - anthropic.claude-opus-5: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-8: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-7: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-6: - input_per_1k: 0.005 - output_per_1k: 0.025 - cache_read_per_1k: 0.0005 - cache_creation_per_1k: 0.00625 - anthropic.claude-opus-4-1: - input_per_1k: 0.015 - output_per_1k: 0.075 - cache_read_per_1k: 0.0015 - cache_creation_per_1k: 0.01875 - anthropic.claude-sonnet-4-6: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - anthropic.claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 - anthropic.claude-haiku-4-5: - input_per_1k: 0.001 - output_per_1k: 0.005 - cache_read_per_1k: 0.0001 - cache_creation_per_1k: 0.00125 - meta.llama3-3-70b-instruct: - input_per_1k: 0.00072 - output_per_1k: 0.00072 - amazon.nova-2-lite: - input_per_1k: 0.0003 - output_per_1k: 0.0025 - amazon.nova-pro: - input_per_1k: 0.0008 - output_per_1k: 0.0032 - amazon.nova-lite: + mistral-large-latest: + input_per_1k: 0.0005 + output_per_1k: 0.0015 + mistral-medium-3-5: + input_per_1k: 0.0015 + output_per_1k: 0.0075 + mistral-medium-latest: + input_per_1k: 0.0004 + output_per_1k: 0.002 + mistral-small-latest: input_per_1k: 0.00006 - output_per_1k: 0.00024 - amazon.nova-micro: - input_per_1k: 0.000035 - output_per_1k: 0.00014 + output_per_1k: 0.00018 + o4-mini: + input_per_1k: 0.0011 + output_per_1k: 0.0044 + cached_input_per_1k: 0.000275 + text-embedding-3-large: + input_per_1k: 0.00013 + output_per_1k: 0 + text-embedding-3-small: + input_per_1k: 0.00002 + output_per_1k: 0 diff --git a/management/internals/modules/agentnetwork/pricing/defaults_test.go b/management/internals/modules/agentnetwork/pricing/defaults_test.go new file mode 100644 index 000000000..99c965687 --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/defaults_test.go @@ -0,0 +1,148 @@ +package pricing + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" +) + +// TestDefaultTable_CoversEveryCatalogModel replaces the proxy's old +// hand-maintained coverage list: because the table is built FROM the +// catalog, drift is impossible by construction — this test guards the +// fold itself (every catalog model of every surfaced provider resolves, +// with exactly the catalog's rates). +func TestDefaultTable_CoversEveryCatalogModel(t *testing.T) { + table := DefaultTable() + for _, p := range catalog.All() { + if len(p.PricingSurfaces) == 0 { + assert.Empty(t, p.Models, "catalog entry %s declares models but no pricing surfaces — those models would never be priced", p.ID) + continue + } + for _, surface := range p.PricingSurfaces { + byModel, ok := table[surface] + require.True(t, ok, "surface %q (provider %s) missing from default table", surface, p.ID) + for _, m := range p.Models { + e, ok := byModel[m.ID] + require.True(t, ok, "%s/%s (provider %s) missing from default table", surface, m.ID, p.ID) + assert.Equal(t, m.InputPer1k, e.InputPer1k, "%s/%s input rate", surface, m.ID) + assert.Equal(t, m.OutputPer1k, e.OutputPer1k, "%s/%s output rate", surface, m.ID) + } + } + } +} + +// TestDefaultTable_NoConflictingContributions enforces the collision rule +// documented on catalog.Provider.PricingSurfaces: when two catalog +// providers contribute the same (surface, model) pair — azure/vertex +// mirroring openai/anthropic, kimi on both surfaces — their rates must be +// identical, because the surface-keyed table can only hold one entry. +// If a provider ever diverges (e.g. Azure reprices a model), this fails +// and the divergence must move to per-provider-record pricing. +func TestDefaultTable_NoConflictingContributions(t *testing.T) { + type contribution struct { + providerID string + entry Entry + } + seen := map[string]map[string]contribution{} + for _, p := range catalog.All() { + for _, surface := range p.PricingSurfaces { + if seen[surface] == nil { + seen[surface] = map[string]contribution{} + } + for _, m := range p.Models { + e := entryFromCatalogModel(m) + if prev, dup := seen[surface][m.ID]; dup { + assert.Equal(t, prev.entry, e, + "%s/%s: %s and %s contribute different rates", surface, m.ID, prev.providerID, p.ID) + continue + } + seen[surface][m.ID] = contribution{providerID: p.ID, entry: e} + } + } + } + // Supplemental entries must never shadow a catalog-contributed model — + // they exist precisely because the catalog does NOT list them. + for surface, models := range supplementalDefaults { + for id := range models { + _, fromCatalog := seen[surface][id] + assert.False(t, fromCatalog, "supplemental %s/%s is now in the catalog — delete the supplemental row", surface, id) + } + } +} + +// TestDefaultTable_AllRatesFiniteNonNegative mirrors the proxy-side +// NewTable validation so a bad catalog edit is caught here, at unit-test +// time, rather than as a chain-build failure on every proxy. +func TestDefaultTable_AllRatesFiniteNonNegative(t *testing.T) { + for surface, models := range DefaultTable() { + for id, e := range models { + for field, v := range map[string]float64{ + "input": e.InputPer1k, + "output": e.OutputPer1k, + "cached_input": e.CachedInputPer1k, + "cache_read": e.CacheReadPer1k, + "cache_creation": e.CacheCreationPer1k, + } { + assert.False(t, v < 0 || math.IsNaN(v) || math.IsInf(v, 0), + "%s/%s: %s rate %v must be finite and non-negative", surface, id, field, v) + } + } + } +} + +// TestDefaultTable_PinnedRates pins rates that previously drifted or are +// easy to mis-enter (carried over from the proxy's retired +// defaults_coverage_test), plus the supplemental entries. +func TestDefaultTable_PinnedRates(t *testing.T) { + table := DefaultTable() + + gpt54 := table["openai"]["gpt-5.4"] + assert.InDelta(t, 0.0025, gpt54.InputPer1k, 1e-9, "gpt-5.4 input") + assert.InDelta(t, 0.015, gpt54.OutputPer1k, 1e-9, "gpt-5.4 output") + assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "gpt-5.4 cached input") + + sonnet := table["bedrock"]["anthropic.claude-sonnet-4-5"] + assert.InDelta(t, 0.003, sonnet.InputPer1k, 1e-9, "bedrock sonnet-4-5 input") + assert.InDelta(t, 0.015, sonnet.OutputPer1k, 1e-9, "bedrock sonnet-4-5 output") + assert.InDelta(t, 0.0003, sonnet.CacheReadPer1k, 1e-9, "bedrock sonnet-4-5 cache read") + assert.InDelta(t, 0.00375, sonnet.CacheCreationPer1k, 1e-9, "bedrock sonnet-4-5 cache creation") + + // Vertex Claude prices under "anthropic" with the bare id. + fable := table["anthropic"]["claude-fable-5"] + assert.InDelta(t, 0.010, fable.InputPer1k, 1e-9, "claude-fable-5 input") + assert.InDelta(t, 0.0125, fable.CacheCreationPer1k, 1e-9, "claude-fable-5 cache creation") + + // Supplementals present on their surfaces. + for surface, ids := range map[string][]string{ + "openai": {"gpt-5", "gpt-5-mini", "gpt-5-nano"}, + "anthropic": {"claude-opus-5", "kimi-k3[1m]", "kimi-k3"}, + "bedrock": {"anthropic.claude-opus-5"}, + } { + for _, id := range ids { + _, ok := table[surface][id] + assert.True(t, ok, "%s/%s must be priced", surface, id) + } + } + + // Embeddings bill input-only — output stays zero. + emb := table["openai"]["text-embedding-3-large"] + assert.Zero(t, emb.OutputPer1k, "embedding output rate must be zero") + assert.Positive(t, emb.InputPer1k, "embedding input rate must be set") +} + +func TestLookupDefault_SurfaceOrder(t *testing.T) { + // kimi-k3 exists on both surfaces; first surface in the slice wins. + e, ok := LookupDefault([]string{"openai", "anthropic"}, "kimi-k3") + require.True(t, ok) + assert.InDelta(t, 0.003, e.InputPer1k, 1e-9) + + _, ok = LookupDefault([]string{"bedrock"}, "gpt-4o") + assert.False(t, ok, "gpt-4o is not a bedrock model") + + _, ok = LookupDefault(nil, "gpt-4o") + assert.False(t, ok, "no surfaces, no match") +} diff --git a/management/internals/modules/agentnetwork/pricing/exampleyaml.go b/management/internals/modules/agentnetwork/pricing/exampleyaml.go new file mode 100644 index 000000000..6217928d6 --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/exampleyaml.go @@ -0,0 +1,109 @@ +package pricing + +import ( + "bytes" + "fmt" + "sort" + "strconv" +) + +// MarshalDefaultsYAML renders the built-in default pricing table (catalog +// + supplementals, WITHOUT any operator override) as the YAML schema +// LoadOverrideFile consumes. It backs the generated +// defaults_llm_pricing.example.yaml so operators start from a file that +// matches the compiled-in rates exactly; a golden test keeps the two in +// sync. Output is deterministic (sorted surfaces and models). +func MarshalDefaultsYAML() []byte { + var b bytes.Buffer + b.WriteString(exampleHeader) + + table := buildDefaultTable() + surfaces := make([]string, 0, len(table)) + for s := range table { + surfaces = append(surfaces, s) + } + sort.Strings(surfaces) + + for _, surface := range surfaces { + fmt.Fprintf(&b, "\n%s:\n", surface) + models := table[surface] + ids := make([]string, 0, len(models)) + for id := range models { + ids = append(ids, id) + } + sort.Strings(ids) + for _, id := range ids { + e := models[id] + fmt.Fprintf(&b, " %s:\n", yamlKey(id)) + fmt.Fprintf(&b, " input_per_1k: %s\n", rate(e.InputPer1k)) + fmt.Fprintf(&b, " output_per_1k: %s\n", rate(e.OutputPer1k)) + if e.CachedInputPer1k > 0 { + fmt.Fprintf(&b, " cached_input_per_1k: %s\n", rate(e.CachedInputPer1k)) + } + if e.CacheReadPer1k > 0 { + fmt.Fprintf(&b, " cache_read_per_1k: %s\n", rate(e.CacheReadPer1k)) + } + if e.CacheCreationPer1k > 0 { + fmt.Fprintf(&b, " cache_creation_per_1k: %s\n", rate(e.CacheCreationPer1k)) + } + } + } + return b.Bytes() +} + +// rate renders a USD-per-1k rate without float noise ("0.00015", not +// "0.000150000000..."). +func rate(v float64) string { + return strconv.FormatFloat(v, 'f', -1, 64) +} + +// yamlKey quotes model ids that YAML would otherwise misparse (e.g. +// "kimi-k3[1m]" starts a flow sequence unquoted). +func yamlKey(id string) string { + for _, r := range id { + switch r { + case '[', ']', '{', '}', ':', '#', ',', '&', '*', '!', '|', '>', '\'', '"', '%', '@', '`': + return strconv.Quote(id) + } + } + return id +} + +const exampleHeader = `# Default LLM pricing used by NetBird's Agent Network cost metering. +# GENERATED from the management catalog — do not edit this file in the +# repository; regenerate with: +# +# go generate ./management/internals/modules/agentnetwork/pricing +# +# Operators: copy this file to /defaults_llm_pricing.yaml (or +# any path configured via management.json: +# +# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } } +# +# ) and adjust the entries you want to change. Management re-reads the +# file periodically (mtime poll, every minute): the live table feeds the +# proxies' cost metering and the dashboard's model-price prefill, so +# edits apply without a restart. Your file only needs the entries you +# want to change — but each entry REPLACES the built-in entry for that +# surface+model whole, so repeat the cache rates you want to keep. +# Unknown fields and negative or non-finite rates are rejected: at +# startup that fails boot (for an explicitly configured path); at +# runtime the previous table is kept and a warning is logged. Deleting +# the file reverts to the built-in defaults below. +# +# Top-level keys are pricing surfaces — the parser shape requests are +# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible +# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock" +# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be +# the normalized id the proxy meters (version/region suffixes stripped). +# +# Values are USD per 1_000 tokens. Optional cache fields: +# cached_input_per_1k OpenAI shape: rate for cached prompt tokens +# (a SUBSET of input tokens). Absent -> cached +# portion bills at input_per_1k. +# cache_read_per_1k Anthropic shape: rate for cache_read tokens +# (ADDITIVE to input). Absent -> input rate. +# cache_creation_per_1k Anthropic shape: rate for cache_creation +# tokens (ADDITIVE to input). Absent -> input +# rate. +` diff --git a/management/internals/modules/agentnetwork/pricing/gen.go b/management/internals/modules/agentnetwork/pricing/gen.go new file mode 100644 index 000000000..c799ab87c --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/gen.go @@ -0,0 +1,20 @@ +//go:build ignore + +// Regenerates defaults_llm_pricing.example.yaml from the compiled-in +// default pricing table. Run via: +// +// go generate ./management/internals/modules/agentnetwork/pricing +package main + +import ( + "log" + "os" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" +) + +func main() { + if err := os.WriteFile("defaults_llm_pricing.example.yaml", pricing.MarshalDefaultsYAML(), 0o644); err != nil { + log.Fatalf("write defaults_llm_pricing.example.yaml: %v", err) + } +} diff --git a/management/internals/modules/agentnetwork/pricing/override.go b/management/internals/modules/agentnetwork/pricing/override.go new file mode 100644 index 000000000..07677dbaa --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/override.go @@ -0,0 +1,249 @@ +package pricing + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "io/fs" + "math" + "os" + "sync" + "time" + + log "github.com/sirupsen/logrus" + "gopkg.in/yaml.v3" +) + +// DefaultFileName is the basename probed under management's datadir when +// AgentNetwork.PricingDefaultsFile doesn't configure an explicit path. +const DefaultFileName = "defaults_llm_pricing.yaml" + +// ReloadInterval is the cadence at which the pricing file's mtime is +// polled for changes. +const ReloadInterval = time.Minute + +// maxFileBytes bounds the pricing file read so a misconfigured path +// (pointed at a huge file) cannot exhaust process memory. +const maxFileBytes = 4 << 20 + +// pricingFile mirrors the on-disk YAML schema — the same schema the +// proxy's retired embedded defaults_pricing.yaml used, so files written +// for it keep working. Keys are pricing surfaces ("openai", "anthropic", +// "bedrock"); nested keys are normalized model ids. +type pricingFile map[string]map[string]struct { + InputPer1k float64 `yaml:"input_per_1k"` + OutputPer1k float64 `yaml:"output_per_1k"` + CachedInputPer1k float64 `yaml:"cached_input_per_1k"` + CacheReadPer1k float64 `yaml:"cache_read_per_1k"` + CacheCreationPer1k float64 `yaml:"cache_creation_per_1k"` +} + +// fileState tracks the watched pricing file across reloads. +var fileState struct { + mu sync.Mutex + path string + mtime int64 +} + +// LoadFile loads the management-side pricing defaults file and makes it +// the live table (merged entry-whole over the compiled-in fallback; see +// DefaultTable). The path stays registered for the periodic reloader, so +// later edits — or the file (re)appearing after deletion — are picked up +// without a restart. +// +// required governs the missing-file case: true for an explicitly +// configured path (a typo must fail startup rather than silently bill +// with built-ins the operator believes they replaced), false for the +// conventional /defaults_llm_pricing.yaml probe (absent file = +// compiled-in defaults, still watched in case it appears). A file that +// exists but is malformed is always an error at load time. +func LoadFile(path string, required bool) error { + if path == "" { + return nil + } + fileState.mu.Lock() + fileState.path = path + fileState.mu.Unlock() + + table, mtime, err := readFile(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) && !required { + log.Infof("agent-network pricing defaults file %s not present; serving built-in defaults", path) + return nil + } + return err + } + storeFileTable(table, mtime) + log.Infof("agent-network pricing defaults loaded from %s", path) + return nil +} + +// StartReloader launches the periodic mtime-poll goroutine for the file +// registered by LoadFile. Runtime failures are lenient — a parse error +// keeps the previously loaded table and logs a warning; a deleted file +// reverts to the compiled-in defaults — so a mid-edit save can never +// take pricing down. Returns immediately when no path was registered. +func StartReloader(ctx context.Context, interval time.Duration) { + fileState.mu.Lock() + path := fileState.path + fileState.mu.Unlock() + if path == "" { + return + } + if interval <= 0 { + interval = ReloadInterval + } + go func() { + t := time.NewTicker(interval) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + reload() + } + } + }() +} + +// reload performs one mtime check + reload cycle. +func reload() { + fileState.mu.Lock() + path, lastMtime := fileState.path, fileState.mtime + fileState.mu.Unlock() + + log.Debugf("agent-network pricing defaults reload: checking %s for changes", path) + + st, err := os.Stat(path) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + // File removed (or not yet created): serve compiled-in + // defaults and reset mtime so a future (re)appearance loads. + if mergedTable.Swap(nil) != nil { + log.Warnf("agent-network pricing defaults file %s removed; reverting to built-in defaults", path) + } + setMtime(0) + return + } + log.Warnf("agent-network pricing defaults reload: stat %s: %v", path, err) + return + } + if st.ModTime().UnixNano() == lastMtime { + log.Debugf("agent-network pricing defaults %s unchanged since last check", path) + return + } + + table, mtime, err := readFile(path) + if err != nil { + // Keep the previously loaded table — never blank prices because + // an operator saved mid-edit. + log.Warnf("agent-network pricing defaults reload failed for %s (keeping previous table): %v", path, err) + return + } + storeFileTable(table, mtime) + log.Infof("agent-network pricing defaults reloaded from %s", path) +} + +func readFile(path string) (map[string]map[string]Entry, int64, error) { + f, err := os.Open(path) + if err != nil { + return nil, 0, fmt.Errorf("open pricing defaults %s: %w", path, err) + } + defer func() { _ = f.Close() }() + + st, err := f.Stat() + if err != nil { + return nil, 0, fmt.Errorf("stat pricing defaults %s: %w", path, err) + } + data, err := io.ReadAll(io.LimitReader(f, maxFileBytes+1)) + if err != nil { + return nil, 0, fmt.Errorf("read pricing defaults %s: %w", path, err) + } + if len(data) > maxFileBytes { + return nil, 0, fmt.Errorf("pricing defaults %s exceeds %d bytes", path, maxFileBytes) + } + table, err := parsePricingYAML(data) + if err != nil { + return nil, 0, fmt.Errorf("parse pricing defaults %s: %w", path, err) + } + return table, st.ModTime().UnixNano(), nil +} + +// storeFileTable merges the parsed file over the compiled-in base and +// publishes the result as the live table. File entries replace the +// built-in entry for the same (surface, model) whole — they are not +// field-merged — and surfaces/models the file doesn't mention keep the +// built-in rates, so a partial file only needs the entries it changes. +func storeFileTable(table map[string]map[string]Entry, mtime int64) { + base := compiledBase() + merged := make(map[string]map[string]Entry, len(base)+len(table)) + for surface, models := range base { + inner := make(map[string]Entry, len(models)) + for id, e := range models { + inner[id] = e + } + merged[surface] = inner + } + for surface, models := range table { + inner, ok := merged[surface] + if !ok { + inner = make(map[string]Entry, len(models)) + merged[surface] = inner + } + for id, e := range models { + inner[id] = e + } + } + mergedTable.Store(&merged) + setMtime(mtime) +} + +func setMtime(v int64) { + fileState.mu.Lock() + fileState.mtime = v + fileState.mu.Unlock() +} + +// parsePricingYAML decodes and validates the pricing YAML. Unknown +// fields are rejected (typos surface instead of silently pricing at 0) +// and every rate must be a finite, non-negative USD amount — the same +// constraints the HTTP API enforces on operator per-provider prices. +func parsePricingYAML(data []byte) (map[string]map[string]Entry, error) { + dec := yaml.NewDecoder(bytes.NewReader(data)) + dec.KnownFields(true) + + var raw pricingFile + if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) { + return nil, fmt.Errorf("decode yaml: %w", err) + } + + out := make(map[string]map[string]Entry, len(raw)) + for surface, models := range raw { + inner := make(map[string]Entry, len(models)) + for model, e := range models { + for field, v := range map[string]float64{ + "input_per_1k": e.InputPer1k, + "output_per_1k": e.OutputPer1k, + "cached_input_per_1k": e.CachedInputPer1k, + "cache_read_per_1k": e.CacheReadPer1k, + "cache_creation_per_1k": e.CacheCreationPer1k, + } { + if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) { + return nil, fmt.Errorf("%s/%s: %s must be a finite, non-negative rate, got %v", surface, model, field, v) + } + } + inner[model] = Entry{ + InputPer1k: e.InputPer1k, + OutputPer1k: e.OutputPer1k, + CachedInputPer1k: e.CachedInputPer1k, + CacheReadPer1k: e.CacheReadPer1k, + CacheCreationPer1k: e.CacheCreationPer1k, + } + } + out[surface] = inner + } + return out, nil +} diff --git a/management/internals/modules/agentnetwork/pricing/override_test.go b/management/internals/modules/agentnetwork/pricing/override_test.go new file mode 100644 index 000000000..86031fc58 --- /dev/null +++ b/management/internals/modules/agentnetwork/pricing/override_test.go @@ -0,0 +1,173 @@ +package pricing + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// resetFileState snapshots and restores the package-level file state so +// tests stay order-independent. +func resetFileState(t *testing.T) { + t.Helper() + prevMerged := mergedTable.Load() + fileState.mu.Lock() + prevPath, prevMtime := fileState.path, fileState.mtime + fileState.mu.Unlock() + t.Cleanup(func() { + mergedTable.Store(prevMerged) + fileState.mu.Lock() + fileState.path, fileState.mtime = prevPath, prevMtime + fileState.mu.Unlock() + }) +} + +func writePricing(t *testing.T, path, yml string) { + t.Helper() + require.NoError(t, os.WriteFile(path, []byte(yml), 0o600)) +} + +func TestLoadFile_MergesOverCompiledDefaults(t *testing.T) { + resetFileState(t) + path := filepath.Join(t.TempDir(), DefaultFileName) + writePricing(t, path, ` +openai: + # Reprice a built-in model. The entry replaces the built-in WHOLE: + # omitting the cache rate here drops the built-in 0.00125 discount. + gpt-4o: + input_per_1k: 0.9 + output_per_1k: 1.8 + # A model NetBird doesn't know at all. + my-private-ft: + input_per_1k: 0.01 + output_per_1k: 0.02 + cached_input_per_1k: 0.005 +gemini: + gemini-pro: + input_per_1k: 0.00125 + output_per_1k: 0.005 +`) + require.NoError(t, LoadFile(path, true)) + table := DefaultTable() + + gpt4o := table["openai"]["gpt-4o"] + assert.InDelta(t, 0.9, gpt4o.InputPer1k, 1e-9, "file rate replaces the compiled-in rate") + assert.Zero(t, gpt4o.CachedInputPer1k, "entries replace whole — omitted cache rate is dropped, not inherited") + + ft := table["openai"]["my-private-ft"] + assert.InDelta(t, 0.005, ft.CachedInputPer1k, 1e-9, "unknown models are added to the surface") + _, ok := table["gemini"]["gemini-pro"] + assert.True(t, ok, "a surface the catalog doesn't declare can be added") + + // Untouched entries keep compiled-in rates (catalog, other surface, + // supplemental). + assert.InDelta(t, 0.00015, table["openai"]["gpt-4o-mini"].InputPer1k, 1e-9, "unlisted model keeps compiled rate") + assert.InDelta(t, 0.003, table["anthropic"]["claude-sonnet-4-5"].InputPer1k, 1e-9, "unlisted surface untouched") + assert.InDelta(t, 0.00125, table["openai"]["gpt-5"].InputPer1k, 1e-9, "supplemental entries untouched") + + // The synthesizer-facing lookup reads the live table too. + e, ok := LookupDefault([]string{"openai"}, "gpt-4o") + require.True(t, ok) + assert.InDelta(t, 0.9, e.InputPer1k, 1e-9, "LookupDefault serves the file-backed rate") +} + +func TestLoadFile_MissingPath(t *testing.T) { + resetFileState(t) + missing := filepath.Join(t.TempDir(), DefaultFileName) + + require.Error(t, LoadFile(missing, true), + "explicitly configured path that doesn't exist must fail startup") + + require.NoError(t, LoadFile(missing, false), + "conventional datadir probe tolerates an absent file (compiled-in defaults serve)") + assert.Nil(t, mergedTable.Load(), "no file, no merged table") + fileState.mu.Lock() + path := fileState.path + fileState.mu.Unlock() + assert.Equal(t, missing, path, "the path stays registered so the reloader picks the file up when it appears") +} + +func TestLoadFile_RejectsInvalid(t *testing.T) { + resetFileState(t) + dir := t.TempDir() + cases := map[string]string{ + "unknown field (typo)": "openai:\n gpt-4o:\n input_per1k: 0.1\n", + "negative rate": "openai:\n gpt-4o:\n input_per_1k: -0.1\n", + "non-numeric rate": "openai:\n gpt-4o:\n input_per_1k: cheap\n", + "not a mapping": "- just\n- a\n- list\n", + } + for name, yml := range cases { + path := filepath.Join(dir, DefaultFileName) + writePricing(t, path, yml) + assert.Error(t, LoadFile(path, true), "case %q must be rejected", name) + } +} + +// TestReload_LifeCycle drives the reloader's single-shot reload through +// its full lifecycle: file edit picked up on mtime change, a broken save +// keeps the previous table, and file removal reverts to the compiled-in +// defaults (then a re-created file loads again). +func TestReload_LifeCycle(t *testing.T) { + resetFileState(t) + path := filepath.Join(t.TempDir(), DefaultFileName) + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.5\n output_per_1k: 1\n") + require.NoError(t, LoadFile(path, true)) + require.InDelta(t, 0.5, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9) + + // Edit: new mtime, new rates. + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.7\n output_per_1k: 1.4\n") + bumpMtime(t, path) + reload() + assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, "edit must be picked up") + + // Broken save: previous table survives. + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: -1\n") + bumpMtime(t, path) + reload() + assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, + "a malformed save must keep the previously loaded table, never blank prices") + + // Removal: compiled-in defaults serve again. + require.NoError(t, os.Remove(path)) + reload() + assert.InDelta(t, 0.0025, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, + "file removal reverts to the compiled-in rate") + + // Re-created file loads without a restart. + writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.9\n output_per_1k: 1.8\n") + bumpMtime(t, path) + reload() + assert.InDelta(t, 0.9, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, + "a file appearing after removal (or after a missing-probe boot) must load") +} + +// bumpMtime pushes the file's mtime forward past the previously recorded +// value — timestamps can otherwise collide within the test's timescale. +func bumpMtime(t *testing.T, path string) { + t.Helper() + st, err := os.Stat(path) + require.NoError(t, err) + next := st.ModTime().Add(2 * 1e9) + require.NoError(t, os.Chtimes(path, next, next)) +} + +// TestExampleYAML_InSyncWithBuiltins is the golden guard for +// defaults_llm_pricing.example.yaml: the shipped example must stay +// byte-identical to what the compiled-in table renders (catalog edits +// require `go generate ./management/internals/modules/agentnetwork/pricing`) +// and must round-trip through the same parser operators' files go +// through, reproducing the compiled-in table exactly. +func TestExampleYAML_InSyncWithBuiltins(t *testing.T) { + onDisk, err := os.ReadFile("defaults_llm_pricing.example.yaml") + require.NoError(t, err, "example file must exist next to the package") + require.Equal(t, string(MarshalDefaultsYAML()), string(onDisk), + "defaults_llm_pricing.example.yaml is stale — run: go generate ./management/internals/modules/agentnetwork/pricing") + + parsed, err := parsePricingYAML(onDisk) + require.NoError(t, err, "the example must be a valid pricing defaults file") + assert.Equal(t, buildDefaultTable(), parsed, + "parsing the example must reproduce the compiled-in table exactly") +} diff --git a/management/internals/modules/agentnetwork/synthesizer.go b/management/internals/modules/agentnetwork/synthesizer.go index 169bdd4fd..64711387b 100644 --- a/management/internals/modules/agentnetwork/synthesizer.go +++ b/management/internals/modules/agentnetwork/synthesizer.go @@ -233,6 +233,11 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([ return nil, err } + costMeterJSON, err := buildCostMeterConfigJSON(enabledProviders, groupIndex) + if err != nil { + return nil, err + } + mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID) applyAccountCollectionControls(&mergedGuardrails, settings) // The proxy guardrail is a per-provider fail-closed backstop; the @@ -248,7 +253,7 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([ // Use the merged decision (account settings OR policy-required redaction), // not the raw account flag, so a policy that mandates PII redaction is // honored by the capture parsers even when the account toggle is off. - middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled) + middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled) priv, pub, err := pickServiceSessionKeys(enabledProviders) if err != nil { @@ -700,7 +705,7 @@ func buildIdentityExtraHeaders(p *types.Provider, extras []catalog.ExtraHeader) // requests bound for gateways like LiteLLM that key budgets and // attribution off request headers. CanMutate is required so its // HeadersAdd / HeadersRemove pass the framework's mutation gate. -func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig { +func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig { // Both parsers receive an explicit capture flag derived from the account's // enable_prompt_collection toggle; nil/unset would default to the legacy // "always emit" behavior in the middleware, which is precisely what we @@ -769,10 +774,13 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt ConfigJSON: []byte("{}"), }, { + // Carries the full pricing table (defaults + per-provider + // operator prices) so the proxy bills without an embedded + // price list; see buildCostMeterConfigJSON. ID: middlewareIDCostMeter, Enabled: true, Slot: rpservice.MiddlewareSlotOnResponse, - ConfigJSON: []byte("{}"), + ConfigJSON: costMeterJSON, }, { ID: middlewareIDLLMResponseParser, diff --git a/management/internals/modules/agentnetwork/synthesizer_pricing.go b/management/internals/modules/agentnetwork/synthesizer_pricing.go new file mode 100644 index 000000000..6d57ffaef --- /dev/null +++ b/management/internals/modules/agentnetwork/synthesizer_pricing.go @@ -0,0 +1,131 @@ +package agentnetwork + +import ( + "encoding/json" + "fmt" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + sharedllm "github.com/netbirdio/netbird/shared/llm" +) + +// costMeterConfig is the JSON shape the proxy-side cost_meter middleware +// expects (mirror-type pattern, same as routerConfig). The top-level +// "pricing" wrapper is the feature-detection signal: an old proxy's config +// struct ignores it as an unknown field, and a new proxy treats its +// absence as "old management" (skips every cost computation and warns). +type costMeterConfig struct { + Pricing *costMeterPricing `json:"pricing,omitempty"` +} + +// costMeterPricing carries the full pricing table: +// - Defaults: surface ("openai"/"anthropic"/"bedrock") -> normalized +// model id -> rates. The full default table ships to every account — +// it is small (~10 KB) and keeps gateway-style providers (which +// enumerate no models) priced for every catalog model. +// - Providers: provider record id (matched against the +// llm.resolved_provider_id metadata llm_router stamps) -> normalized +// model id -> rates. Entries are fully materialized here at synth +// time — default cache rates already folded in — so the proxy does +// two map lookups and no merging. +type costMeterPricing struct { + Defaults map[string]map[string]pricing.Entry `json:"defaults,omitempty"` + Providers map[string]map[string]pricing.Entry `json:"providers,omitempty"` +} + +// buildCostMeterConfigJSON assembles the cost_meter middleware config +// from the default pricing table plus the operator's stored per-provider +// model prices. Same orphan rule as the router: a provider no enabled +// policy authorises is unreachable, so its prices are not shipped. +// +// Overlay semantics per model row: +// - The row's model id is normalized exactly the way the proxy's +// request parser normalizes the ids it meters (bedrock ARN/region/ +// version stripping, vertex "@version" stripping), so the per-record +// lookup key compares equal to llm.model at billing time. +// - The entry starts from the default entry for that model (when one +// exists) to inherit cache rates the operator didn't state. +// - Operator input/output overlay verbatim — including an explicit 0, +// which prices the model as free (self-hosted / internal endpoints) +// rather than silently reverting to list price. +// - Cache-rate pointers overlay only when non-nil: nil means "inherit +// the default", an explicit 0 means "no discount, bill this bucket +// at the input rate". +func buildCostMeterConfigJSON(providers []*types.Provider, groupIndex map[string][]string) ([]byte, error) { + cfg := costMeterConfig{Pricing: &costMeterPricing{ + Defaults: pricing.DefaultTable(), + }} + + perRecord := make(map[string]map[string]pricing.Entry) + for _, p := range providers { + if _, hasPolicy := groupIndex[p.ID]; !hasPolicy { + // Orphan: unreachable via the router, so unpriceable. + continue + } + if len(p.Models) == 0 { + // Gateway-style "claim every model" provider: the defaults + // table is its price list. + continue + } + entry, _ := catalog.Lookup(p.ProviderID) + models := make(map[string]pricing.Entry, len(p.Models)) + for _, m := range p.Models { + id := normalizePricingModelID(p.ProviderID, m.ID) + if id == "" { + continue + } + if _, dup := models[id]; dup { + // First occurrence wins on post-normalization duplicates, + // matching providerModelIDs' dedup order for routing. + continue + } + models[id] = materializeEntry(entry.PricingSurfaces, id, m) + } + if len(models) > 0 { + perRecord[p.ID] = models + } + } + if len(perRecord) > 0 { + cfg.Pricing.Providers = perRecord + } + + out, err := json.Marshal(cfg) + if err != nil { + return nil, fmt.Errorf("marshal cost_meter middleware config: %w", err) + } + return out, nil +} + +// normalizePricingModelID maps an operator-entered model id onto the +// normalized id the proxy's request parser emits as llm.model — the key +// the cost meter looks up at billing time. +func normalizePricingModelID(catalogProviderID, modelID string) string { + switch { + case catalog.IsBedrockPathStyle(catalogProviderID): + return sharedllm.NormalizeBedrockModel(modelID) + case catalog.IsVertexPathStyle(catalogProviderID): + return sharedllm.NormalizeVertexModel(modelID) + default: + return modelID + } +} + +// materializeEntry folds the default entry for (surfaces, model) — when +// one exists — under the operator's stored prices, producing the fully +// materialized wire entry. +func materializeEntry(surfaces []string, normalizedID string, m types.ProviderModel) pricing.Entry { + e, _ := pricing.LookupDefault(surfaces, normalizedID) // zero Entry on miss + e.InputPer1k = m.InputPer1k + e.OutputPer1k = m.OutputPer1k + if m.CachedInputPer1k != nil { + e.CachedInputPer1k = *m.CachedInputPer1k + } + if m.CacheReadPer1k != nil { + e.CacheReadPer1k = *m.CacheReadPer1k + } + if m.CacheCreationPer1k != nil { + e.CacheCreationPer1k = *m.CacheCreationPer1k + } + return e +} diff --git a/management/internals/modules/agentnetwork/synthesizer_pricing_test.go b/management/internals/modules/agentnetwork/synthesizer_pricing_test.go new file mode 100644 index 000000000..83961878a --- /dev/null +++ b/management/internals/modules/agentnetwork/synthesizer_pricing_test.go @@ -0,0 +1,105 @@ +package agentnetwork + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" +) + +func fptr(v float64) *float64 { return &v } + +func decodeCostMeterConfig(t *testing.T, raw []byte) costMeterConfig { + t.Helper() + var cfg costMeterConfig + require.NoError(t, json.Unmarshal(raw, &cfg), "cost meter config must round-trip") + require.NotNil(t, cfg.Pricing, "pricing wrapper must be present") + return cfg +} + +// TestBuildCostMeterConfig_BedrockModelNormalization: the operator may +// paste region-prefixed, versioned, or ARN-wrapped Bedrock ids; the +// per-record map must be keyed by the normalized id the request parser +// emits as llm.model, or the lookup never hits at billing time. +func TestBuildCostMeterConfig_BedrockModelNormalization(t *testing.T) { + bedrock := &types.Provider{ + ID: "prov-bedrock", + ProviderID: "bedrock_api", + Enabled: true, + Models: []types.ProviderModel{ + {ID: "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 0.0033, OutputPer1k: 0.0165}, + // Post-normalization duplicate of the row above under a + // different regional spelling — first occurrence wins. + {ID: "us.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 9.9, OutputPer1k: 9.9}, + }, + } + raw, err := buildCostMeterConfigJSON([]*types.Provider{bedrock}, map[string][]string{"prov-bedrock": {"grp"}}) + require.NoError(t, err) + cfg := decodeCostMeterConfig(t, raw) + + models := cfg.Pricing.Providers["prov-bedrock"] + require.Len(t, models, 1, "both spellings normalize to one model; first row wins") + e, ok := models["anthropic.claude-sonnet-4-5"] + require.True(t, ok, "key must be the normalized id the parser emits, not the operator's raw spelling") + assert.InDelta(t, 0.0033, e.InputPer1k, 1e-9, "first row's rate wins the dedup") + assert.InDelta(t, 0.0003, e.CacheReadPer1k, 1e-9, "cache read inherited from the bedrock default entry") + assert.InDelta(t, 0.00375, e.CacheCreationPer1k, 1e-9, "cache creation inherited from the bedrock default entry") +} + +// TestBuildCostMeterConfig_CacheRateNilVsZero pins the pointer semantics: +// nil inherits the default cache rate, explicit 0 clears it (that bucket +// bills at the input rate on the proxy). +func TestBuildCostMeterConfig_CacheRateNilVsZero(t *testing.T) { + p := &types.Provider{ + ID: "prov-oai", + ProviderID: "openai_api", + Enabled: true, + Models: []types.ProviderModel{ + {ID: "gpt-4o", InputPer1k: 0.002, OutputPer1k: 0.008}, // nil → inherit 0.00125 + {ID: "gpt-4o-mini", InputPer1k: 0.0001, OutputPer1k: 0.0005, CachedInputPer1k: fptr(0)}, // explicit 0 → no discount + {ID: "my-custom-ft", InputPer1k: 0.01, OutputPer1k: 0.02, CachedInputPer1k: fptr(0.005)}, // unknown model, explicit rate + }, + } + raw, err := buildCostMeterConfigJSON([]*types.Provider{p}, map[string][]string{"prov-oai": {"grp"}}) + require.NoError(t, err) + cfg := decodeCostMeterConfig(t, raw) + models := cfg.Pricing.Providers["prov-oai"] + + assert.InDelta(t, 0.00125, models["gpt-4o"].CachedInputPer1k, 1e-9, "nil cache pointer inherits the default rate") + assert.Zero(t, models["gpt-4o-mini"].CachedInputPer1k, "explicit 0 overrides the default (0.000075) — bucket bills at input rate") + custom := models["my-custom-ft"] + assert.InDelta(t, 0.005, custom.CachedInputPer1k, 1e-9, "unknown model keeps the operator's explicit cache rate") + assert.Zero(t, custom.CacheReadPer1k, "no default to inherit for a model outside the catalog") +} + +// TestBuildCostMeterConfig_OrphanAndGatewayProviders: an orphan (no +// authorising policy) is unreachable so its prices must not ship; a +// gateway with no model rows relies on the defaults table and gets no +// per-record entry. +func TestBuildCostMeterConfig_OrphanAndGatewayProviders(t *testing.T) { + orphan := &types.Provider{ + ID: "prov-orphan", + ProviderID: "openai_api", + Enabled: true, + Models: []types.ProviderModel{{ID: "gpt-4o", InputPer1k: 1, OutputPer1k: 1}}, + } + gateway := &types.Provider{ + ID: "prov-litellm", + ProviderID: "litellm_proxy", + Enabled: true, + Models: []types.ProviderModel{}, + } + raw, err := buildCostMeterConfigJSON( + []*types.Provider{orphan, gateway}, + map[string][]string{"prov-litellm": {"grp"}}, // orphan has no policy + ) + require.NoError(t, err) + cfg := decodeCostMeterConfig(t, raw) + + assert.NotContains(t, cfg.Pricing.Providers, "prov-orphan", "orphan provider prices must not ship") + assert.NotContains(t, cfg.Pricing.Providers, "prov-litellm", "empty-models gateway needs no per-record entry") + assert.NotEmpty(t, cfg.Pricing.Defaults["openai"], "defaults still ship so the gateway's catalog-model traffic is priced") +} diff --git a/management/internals/modules/agentnetwork/synthesizer_test.go b/management/internals/modules/agentnetwork/synthesizer_test.go index 8a18a9b59..7b14f8209 100644 --- a/management/internals/modules/agentnetwork/synthesizer_test.go +++ b/management/internals/modules/agentnetwork/synthesizer_test.go @@ -33,14 +33,17 @@ func newSynthTestSettings() *types.Settings { func newSynthTestProvider() *types.Provider { return &types.Provider{ - ID: "prov-1", - AccountID: testAccountID, - ProviderID: "openai_api", - Name: "OpenAI", - UpstreamURL: "https://api.openai.com", - APIKey: "sk-test-key", - Enabled: true, - Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015}}, + ID: "prov-1", + AccountID: testAccountID, + ProviderID: "openai_api", + Name: "OpenAI", + UpstreamURL: "https://api.openai.com", + APIKey: "sk-test-key", + Enabled: true, + // Prices deliberately differ from the catalog's gpt-5.4 rates + // (0.0025/0.015) so pricing tests can prove the operator's + // stored price overlays the catalog default. + Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}}, SessionPrivateKey: "test-priv-key", SessionPublicKey: "test-pub-key", CreatedAt: time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC), @@ -214,7 +217,27 @@ func TestSynthesizeServices_HappyPath(t *testing.T) { assert.Equal(t, middlewareIDCostMeter, mws[6].ID, "seventh middleware is the cost meter") assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[6].Slot, "cost meter runs on_response") - assert.Equal(t, []byte("{}"), mws[6].ConfigJSON, "cost meter carries an explicit empty config") + + var costCfg costMeterConfig + require.NoError(t, json.Unmarshal(mws[6].ConfigJSON, &costCfg), "cost meter config must unmarshal") + require.NotNil(t, costCfg.Pricing, "cost meter config must carry the pricing table — its absence tells the proxy management predates config-delivered pricing") + + gpt4o, ok := costCfg.Pricing.Defaults["openai"]["gpt-4o"] + require.True(t, ok, "the full default table ships regardless of the account's providers") + assert.InDelta(t, 0.0025, gpt4o.InputPer1k, 1e-9, "default gpt-4o input rate comes from the catalog") + + openaiPrices, ok := costCfg.Pricing.Providers[openai.ID] + require.True(t, ok, "operator-priced provider must have a per-record entry") + gpt54, ok := openaiPrices["gpt-5.4"] + require.True(t, ok, "operator's model row keys the per-record map") + assert.InDelta(t, 0.004, gpt54.InputPer1k, 1e-9, "operator input price overlays the catalog default (0.0025)") + assert.InDelta(t, 0.02, gpt54.OutputPer1k, 1e-9, "operator output price overlays the catalog default (0.015)") + assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "cache rate the operator didn't state is inherited from the default entry") + + opus, ok := costCfg.Pricing.Providers[anthropic.ID]["claude-opus-4-7"] + require.True(t, ok, "anthropic's model row keys its per-record map") + assert.Zero(t, opus.InputPer1k, "operator-stored zero prices ship verbatim — an explicit $0 model bills as free, it does not revert to list price") + assert.InDelta(t, 0.0005, opus.CacheReadPer1k, 1e-9, "cache rates still inherit from the default entry") assert.Equal(t, middlewareIDLLMResponseParser, mws[7].ID, "eighth middleware is the response parser") assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[7].Slot, "response parser runs on_response") diff --git a/management/internals/modules/agentnetwork/types/provider.go b/management/internals/modules/agentnetwork/types/provider.go index b3287168e..96242f45f 100644 --- a/management/internals/modules/agentnetwork/types/provider.go +++ b/management/internals/modules/agentnetwork/types/provider.go @@ -14,10 +14,24 @@ import ( // ProviderModel is one row in the provider's models list. The operator // pins the per-1k input/output price for cost tracking; ID is the // model identifier the upstream provider expects on the wire. +// +// The three cache rates are pointers because absence is meaningful: nil +// means "inherit NetBird's default rate for this model" (folded in at +// synthesis time), while an explicit 0 means "no discount — bill this +// cache bucket at the input rate". type ProviderModel struct { ID string `json:"id"` InputPer1k float64 `json:"input_per_1k"` OutputPer1k float64 `json:"output_per_1k"` + // CachedInputPer1k is the OpenAI-shape rate for cached prompt tokens + // (a subset of input tokens). + CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"` + // CacheReadPer1k is the Anthropic-shape rate for cache-read tokens + // (additive to input tokens). + CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"` + // CacheCreationPer1k is the Anthropic-shape rate for cache-creation + // tokens (additive to input tokens). + CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"` } // Provider is an Agent Network AI provider record persisted per account. @@ -128,9 +142,12 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) { if req.Models != nil { for _, m := range *req.Models { p.Models = append(p.Models, ProviderModel{ - ID: m.Id, - InputPer1k: m.InputPer1k, - OutputPer1k: m.OutputPer1k, + ID: m.Id, + InputPer1k: m.InputPer1k, + OutputPer1k: m.OutputPer1k, + CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k), + CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k), + CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k), }) } } @@ -164,9 +181,12 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { models := make([]api.AgentNetworkProviderModel, 0, len(p.Models)) for _, m := range p.Models { models = append(models, api.AgentNetworkProviderModel{ - Id: m.ID, - InputPer1k: m.InputPer1k, - OutputPer1k: m.OutputPer1k, + Id: m.ID, + InputPer1k: m.InputPer1k, + OutputPer1k: m.OutputPer1k, + CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k), + CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k), + CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k), }) } created := p.CreatedAt @@ -201,11 +221,27 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { return resp } +// copyFloatPtr returns a fresh pointer to the same value, or nil. Keeps +// stored models and API payloads from aliasing each other's rate fields. +func copyFloatPtr(v *float64) *float64 { + if v == nil { + return nil + } + out := *v + return &out +} + // Copy returns a deep copy of the provider. func (p *Provider) Copy() *Provider { clone := *p if p.Models != nil { - clone.Models = append([]ProviderModel(nil), p.Models...) + clone.Models = make([]ProviderModel, len(p.Models)) + for i, m := range p.Models { + m.CachedInputPer1k = copyFloatPtr(m.CachedInputPer1k) + m.CacheReadPer1k = copyFloatPtr(m.CacheReadPer1k) + m.CacheCreationPer1k = copyFloatPtr(m.CacheCreationPer1k) + clone.Models[i] = m + } } if p.ExtraValues != nil { clone.ExtraValues = make(map[string]string, len(p.ExtraValues)) diff --git a/management/internals/modules/agentnetwork/wire_shape_test.go b/management/internals/modules/agentnetwork/wire_shape_test.go index b574ab3e1..779dd77f9 100644 --- a/management/internals/modules/agentnetwork/wire_shape_test.go +++ b/management/internals/modules/agentnetwork/wire_shape_test.go @@ -103,6 +103,12 @@ func TestSynthesizedService_WireShape(t *testing.T) { assert.Equal(t, middlewareIDCostMeter, mws[6].GetId(), "seventh middleware id") assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[6].GetSlot(), "cost meter slot") + var costCfg costMeterConfig + require.NoError(t, json.Unmarshal(mws[6].GetConfigJson(), &costCfg), "cost meter config JSON must decode from the wire") + require.NotNil(t, costCfg.Pricing, "the pricing table must travel on the wire — the proxy has no embedded price list to fall back to") + assert.NotEmpty(t, costCfg.Pricing.Defaults["openai"], "default table rides in every mapping") + assert.NotEmpty(t, costCfg.Pricing.Defaults["anthropic"], "default table covers all surfaces") + assert.NotEmpty(t, costCfg.Pricing.Defaults["bedrock"], "default table covers all surfaces") assert.Equal(t, middlewareIDLLMResponseParser, mws[7].GetId(), "eighth middleware id") assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[7].GetSlot(), "response parser slot") diff --git a/management/internals/server/config/config.go b/management/internals/server/config/config.go index a77d5c19b..dc60ed822 100644 --- a/management/internals/server/config/config.go +++ b/management/internals/server/config/config.go @@ -55,6 +55,8 @@ type Config struct { ReverseProxy ReverseProxy + AgentNetwork AgentNetwork + // disable default all-to-all policy DisableDefaultPolicy bool @@ -185,6 +187,25 @@ type StoreConfig struct { Engine types.Engine } +// AgentNetwork contains agent-network (LLM gateway) configuration. +type AgentNetwork struct { + // PricingDefaultsFile is the path to the YAML file holding the default + // LLM pricing table (defaults_llm_pricing.yaml). A relative path is + // resolved against , so a bare filename lands alongside the + // store. Empty falls back to probing /defaults_llm_pricing.yaml; + // with no file present the compiled-in defaults serve. Schema: surface ("openai"/"anthropic"/ + // "bedrock") -> model -> rates in USD per 1k tokens (input_per_1k, + // output_per_1k, and the optional cached_input_per_1k / + // cache_read_per_1k / cache_creation_per_1k). File entries replace the + // compiled-in entry for the same surface+model whole; everything else + // keeps the compiled-in rates. The file is re-read periodically (mtime + // poll), and the live table feeds both the synthesizer (what proxies + // bill with) and the dashboard's catalog endpoint (what model rows + // prefill with). An explicitly configured path that fails to load + // fails startup; runtime reload errors keep the previous table. + PricingDefaultsFile string +} + // ReverseProxy contains reverse proxy configuration in front of management. type ReverseProxy struct { // TrustedHTTPProxies represents a list of trusted HTTP proxies by their IP prefixes. diff --git a/proxy/internal/llm/bedrock_model.go b/proxy/internal/llm/bedrock_model.go deleted file mode 100644 index a4c4704f7..000000000 --- a/proxy/internal/llm/bedrock_model.go +++ /dev/null @@ -1,38 +0,0 @@ -package llm - -import ( - "regexp" - "strings" -) - -// bedrockRegionPrefixes are the cross-region inference-profile prefixes that -// front a Bedrock model id (e.g. "eu.anthropic.claude-..."). -var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."} - -// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]" -// version/throughput suffix of a Bedrock model id. -var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`) - -// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile -// prefix, and the version/throughput suffix from a Bedrock model id so it -// matches the catalog/pricing key, e.g. -// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5" -// and the inference-profile ARN's last segment likewise. It is the single -// source of truth shared by the request parser (which normalizes the request -// model from the URL path) and the router (which normalizes the operator's -// registered Bedrock model ids so both sides compare equal). -func NormalizeBedrockModel(modelID string) string { - m := modelID - if strings.HasPrefix(m, "arn:") { - if i := strings.LastIndex(m, "/"); i >= 0 { - m = m[i+1:] - } - } - for _, p := range bedrockRegionPrefixes { - if strings.HasPrefix(m, p) { - m = m[len(p):] - break - } - } - return bedrockVersionSuffix.ReplaceAllString(m, "") -} diff --git a/proxy/internal/llm/fixtures/pricing.yaml b/proxy/internal/llm/fixtures/pricing.yaml deleted file mode 100644 index 3d26ff803..000000000 --- a/proxy/internal/llm/fixtures/pricing.yaml +++ /dev/null @@ -1,59 +0,0 @@ -# Realistic-pricing starter for llm_observability. Drop this into the -# directory you point the proxy at via --plugin-data-dir, then reference it -# from the target's plugin config: -# -# plugins: -# - id: llm_observability -# enabled: true -# params: -# pricing_path: pricing.yaml -# -# Values are USD per 1_000 tokens. Public list prices drift; treat this as a -# starting point and keep your production copy current. - -openai: - # GPT-5 family - gpt-5: - input_per_1k: 0.00125 - output_per_1k: 0.01 - gpt-5-mini: - input_per_1k: 0.00025 - output_per_1k: 0.002 - gpt-5-nano: - input_per_1k: 0.00005 - output_per_1k: 0.0004 - gpt-5.4: - input_per_1k: 0.00125 - output_per_1k: 0.01 - # GPT-4o family - gpt-4o: - input_per_1k: 0.0025 - output_per_1k: 0.01 - gpt-4o-mini: - input_per_1k: 0.00015 - output_per_1k: 0.0006 - # Embeddings - text-embedding-3-large: - input_per_1k: 0.00013 - output_per_1k: 0 - text-embedding-3-small: - input_per_1k: 0.00002 - output_per_1k: 0 - -anthropic: - # Claude 4.x family - claude-opus-4-7: - input_per_1k: 0.015 - output_per_1k: 0.075 - claude-sonnet-4-7: - input_per_1k: 0.003 - output_per_1k: 0.015 - claude-sonnet-4-6: - input_per_1k: 0.003 - output_per_1k: 0.015 - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - claude-haiku-4-5: - input_per_1k: 0.0008 - output_per_1k: 0.004 diff --git a/proxy/internal/llm/model.go b/proxy/internal/llm/model.go new file mode 100644 index 000000000..76ccfeccf --- /dev/null +++ b/proxy/internal/llm/model.go @@ -0,0 +1,21 @@ +package llm + +import ( + sharedllm "github.com/netbirdio/netbird/shared/llm" +) + +// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile +// prefix, and the version/throughput suffix from a Bedrock model id so it +// matches the catalog/pricing key. Thin delegate to the shared implementation +// (shared/llm), which management also uses at synthesis time so both sides of +// the pricing / routing contract normalize identically. +func NormalizeBedrockModel(modelID string) string { + return sharedllm.NormalizeBedrockModel(modelID) +} + +// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id +// so it matches the catalog/pricing key. Thin delegate to shared/llm, kept +// beside NormalizeBedrockModel for the same contract reason. +func NormalizeVertexModel(modelID string) string { + return sharedllm.NormalizeVertexModel(modelID) +} diff --git a/proxy/internal/llm/pricing/defaults_coverage_test.go b/proxy/internal/llm/pricing/defaults_coverage_test.go deleted file mode 100644 index 8df1557ea..000000000 --- a/proxy/internal/llm/pricing/defaults_coverage_test.go +++ /dev/null @@ -1,65 +0,0 @@ -package pricing - -import ( - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -// TestDefaultTable_FirstPartyModelCoverage guards the embedded defaults against -// silent drift/gaps: every metered first-party model the management catalog -// enumerates must resolve to a price, and a few rates that previously drifted -// are pinned to their LiteLLM-validated values. Keep this list in step with the -// catalog (management/server/agentnetwork/catalog) when adding models. -func TestDefaultTable_FirstPartyModelCoverage(t *testing.T) { - tbl := DefaultTable() - require.NotNil(t, tbl, "embedded default pricing table must load") - - mustPrice := map[string][]string{ - // openai parser covers openai_api, azure_openai_api, and mistral_api. - "openai": { - "gpt-5.5", "gpt-5.5-pro", "gpt-5.4", "gpt-5.4-mini", "gpt-5.4-nano", - "gpt-5.3-codex", "gpt-5.3-chat-latest", "o4-mini", - "gpt-4.1", "gpt-4.1-mini", "gpt-4.1-nano", "gpt-4o", "gpt-4o-mini", - "gpt-4-turbo", "gpt-3.5-turbo", "gpt-35-turbo", - "text-embedding-3-large", "text-embedding-3-small", - "mistral-large-latest", "mistral-medium-3-5", "codestral-2508", - "ministral-8b-latest", "mistral-embed", - }, - "anthropic": { - "claude-fable-5", "claude-opus-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6", - "claude-opus-4-1", "claude-sonnet-4-6", "claude-sonnet-4-5", "claude-haiku-4-5", - }, - // bedrock keys are the normalized ids the request parser emits. - "bedrock": { - "anthropic.claude-opus-5", "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6", - "anthropic.claude-opus-4-1", "anthropic.claude-sonnet-4-6", "anthropic.claude-sonnet-4-5", - "anthropic.claude-haiku-4-5", "meta.llama3-3-70b-instruct", - "amazon.nova-pro", "amazon.nova-lite", "amazon.nova-micro", "amazon.nova-2-lite", - }, - } - for provider, models := range mustPrice { - for _, m := range models { - _, ok := tbl.Cost(provider, m, 1000, 1000, 0, 0) - assert.True(t, ok, "%s/%s must be priced in the embedded defaults", provider, m) - } - } - - // Pin per-direction rates independently (input-only then output-only) so a - // swap or skew of input<->output that preserves the combined total is still - // caught — these are rates that previously drifted or are easy to mis-enter. - in, ok := tbl.Cost("openai", "gpt-5.4", 1000, 0, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.0025, in, 1e-9, "gpt-5.4 input = 0.0025 per 1k") - out, ok := tbl.Cost("openai", "gpt-5.4", 0, 1000, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.015, out, 1e-9, "gpt-5.4 output = 0.015 per 1k") - - in, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 1000, 0, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.003, in, 1e-9, "bedrock sonnet-4-5 input = 0.003 per 1k") - out, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 0, 1000, 0, 0) - require.True(t, ok) - assert.InDelta(t, 0.015, out, 1e-9, "bedrock sonnet-4-5 output = 0.015 per 1k") -} diff --git a/proxy/internal/llm/pricing/pricing.go b/proxy/internal/llm/pricing/pricing.go index b77000000..ce6e636cf 100644 --- a/proxy/internal/llm/pricing/pricing.go +++ b/proxy/internal/llm/pricing/pricing.go @@ -1,102 +1,30 @@ -// Package pricing implements the embedded-default + override pricing table -// shared by middleware that converts LLM token usage into a USD cost -// estimate. The table is hot-reloadable from a basename under the proxy -// data directory; missing override files keep the embedded defaults so -// cost annotation works without operator action. +// Package pricing implements the pricing table and cost formula the +// cost_meter middleware uses to convert LLM token usage into a USD cost +// estimate. The table's content arrives from the management server inside +// cost_meter's middleware config (synthesized from the catalog plus the +// operator's stored per-provider prices) — the proxy carries no embedded +// price list. Price updates ride the ordinary mapping push: a chain +// rebuild constructs a fresh table, so there is nothing to reload. package pricing import ( - "bytes" - "context" - _ "embed" - "errors" "fmt" - "io" - "io/fs" "math" - "path/filepath" - "regexp" - "strings" - "sync" - "sync/atomic" - "time" - - log "github.com/sirupsen/logrus" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/metric" - "gopkg.in/yaml.v3" ) -//go:embed defaults_pricing.yaml -var defaultPricingYAML []byte - -var ( - defaultTableOnce sync.Once - defaultTablePtr *Table -) - -// DefaultTable returns the pricing table embedded in the binary. The result -// is parsed once and shared; callers must not mutate the returned value. -// Cost annotation works without any operator action because every loader -// starts with this table. -func DefaultTable() *Table { - defaultTableOnce.Do(func() { - t, err := parsePricingBytes(defaultPricingYAML) - if err != nil { - panic(fmt.Sprintf("llmobs: embedded default pricing failed to parse: %v", err)) - } - defaultTablePtr = t - }) - return defaultTablePtr -} - -// mergeOver returns a new Table containing every entry from base, with any -// matching entry from overlay replacing the base value. Either argument may -// be nil. Result is a fresh allocation so callers can mutate / Store safely. -func mergeOver(base, overlay *Table) *Table { - if overlay == nil || len(overlay.entries) == 0 { - return base - } - if base == nil || len(base.entries) == 0 { - return overlay - } - out := make(map[string]map[string]Entry, len(base.entries)) - for provider, models := range base.entries { - inner := make(map[string]Entry, len(models)) - for model, e := range models { - inner[model] = e - } - out[provider] = inner - } - for provider, models := range overlay.entries { - inner, ok := out[provider] - if !ok { - inner = make(map[string]Entry, len(models)) - out[provider] = inner - } - for model, e := range models { - inner[model] = e - } - } - return &Table{entries: out} -} - // Entry is a single model's input and output pricing, expressed in USD per // 1000 tokens. // // CachedInputPer1K applies to OpenAI's cached prompt tokens, which are a // subset of input_tokens — when set, the cached portion is billed at this // rate and the non-cached remainder at InputPer1K. Zero means "no discount -// configured", and cached tokens are billed at InputPer1K (matches current -// behaviour where cached counts weren't extracted at all). +// configured", and cached tokens are billed at InputPer1K. // // CacheReadPer1K and CacheCreationPer1K apply to Anthropic's two prompt- // cache fields, which are additive to input_tokens: cache_read is the // cheaper read-from-cache rate, cache_creation is the more expensive // write-to-cache rate. Zero means "no rate configured" and the -// corresponding token bucket is billed at InputPer1K. This is more -// accurate than today's behaviour, where Anthropic's cache tokens are -// ignored and not charged at all. +// corresponding token bucket is billed at InputPer1K. type Entry struct { InputPer1K float64 OutputPer1K float64 @@ -105,33 +33,102 @@ type Entry struct { CacheCreationPer1K float64 } -// Table is a provider-to-model pricing lookup. Instances are immutable once -// built and are swapped atomically by Loader. +// EntryJSON is the wire shape of a pricing entry inside cost_meter's +// middleware config. Field names are the management→proxy contract; the +// management synthesizer marshals the same names (its pricing.Entry). +type EntryJSON struct { + InputPer1K float64 `json:"input_per_1k"` + OutputPer1K float64 `json:"output_per_1k"` + CachedInputPer1K float64 `json:"cached_input_per_1k"` + CacheReadPer1K float64 `json:"cache_read_per_1k"` + CacheCreationPer1K float64 `json:"cache_creation_per_1k"` +} + +// Table is a provider-surface-to-model pricing lookup. Instances are +// immutable once built; a mapping update builds a whole new middleware +// instance (and with it a new table) rather than mutating this one. type Table struct { entries map[string]map[string]Entry } +// NewEntries validates and converts a wire-shape map (surface-or-record -> +// model -> rates) into the internal representation. Every rate must be a +// finite, non-negative USD amount; a violation is returned as an error so +// a corrupt config fails the chain build loudly instead of mispricing. +// Management validates the same constraints at its API boundary, so this +// is defense-in-depth. Nil input yields an empty (never-matching) map. +func NewEntries(raw map[string]map[string]EntryJSON) (map[string]map[string]Entry, error) { + out := make(map[string]map[string]Entry, len(raw)) + for outer, models := range raw { + inner := make(map[string]Entry, len(models)) + for model, e := range models { + for field, v := range map[string]float64{ + "input_per_1k": e.InputPer1K, + "output_per_1k": e.OutputPer1K, + "cached_input_per_1k": e.CachedInputPer1K, + "cache_read_per_1k": e.CacheReadPer1K, + "cache_creation_per_1k": e.CacheCreationPer1K, + } { + if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) { + return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", outer, model, field, v) + } + } + // EntryJSON and Entry are field-identical (tags aside), so a + // direct conversion carries all five rates. + inner[model] = Entry(e) + } + out[outer] = inner + } + return out, nil +} + +// NewTable builds an immutable Table from the wire-shape defaults map. +// See NewEntries for validation semantics. +func NewTable(raw map[string]map[string]EntryJSON) (*Table, error) { + entries, err := NewEntries(raw) + if err != nil { + return nil, err + } + return &Table{entries: entries}, nil +} + +// Lookup returns the entry for the given provider surface and model. +func (t *Table) Lookup(provider, model string) (Entry, bool) { + if t == nil { + return Entry{}, false + } + byModel, ok := t.entries[provider] + if !ok { + return Entry{}, false + } + e, ok := byModel[model] + return e, ok +} + +// Has reports whether the provider/model pair is present in the table. +func (t *Table) Has(provider, model string) bool { + _, ok := t.Lookup(provider, model) + return ok +} + // Cost returns the estimated USD cost for the given token counts. ok is // false when the provider or model is not present in the table; the caller // can still emit token metrics with a model=unknown label. -// -// Provider-shape semantics for cached / cache-creation counts: -// -// - OpenAI: cachedInput is a SUBSET of inTokens. The cached portion is -// billed at CachedInputPer1K (or InputPer1K when no override), and the -// non-cached remainder of inTokens at InputPer1K. cacheCreation is -// ignored (OpenAI has no analogue). -// - Anthropic: cachedInput (cache_read) and cacheCreation are ADDITIVE to -// inTokens. The three buckets are billed at CacheReadPer1K, -// CacheCreationPer1K, and InputPer1K respectively, each falling back -// to InputPer1K when the corresponding rate is zero. -// - Other providers: cached and cacheCreation are ignored; cost is -// inTokens*InputPer1K + outTokens*OutputPer1K. func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (float64, bool) { c, ok := t.Costs(provider, model, inTokens, outTokens, cachedInput, cacheCreation) return c.TotalUSD, ok } +// Costs returns the estimated USD cost split for the given token counts. +// The provider surface selects the cache formula; see EntryCosts. +func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) { + entry, ok := t.Lookup(provider, model) + if !ok { + return Costs{}, false + } + return EntryCosts(entry, provider, inTokens, outTokens, cachedInput, cacheCreation), true +} + // Costs is a per-request cost split. The four per-bucket fields are the base // of the breakdown — one per token bucket the provider bills separately — and // the two aggregates are derived from them: @@ -165,9 +162,25 @@ func newCosts(input, cachedInput, cacheCreation, output float64) Costs { } } -// Costs returns the estimated USD cost split for the given token counts, with -// the same semantics as Cost. -func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) { +// EntryCosts computes the USD cost split for the given entry and token +// counts. The surface (the llm.provider value the request parser stamps) +// selects the cache formula; the entry may come from the surface-keyed +// defaults table or from a per-provider-record override — the math is +// identical either way. +// +// Provider-shape semantics for cached / cache-creation counts: +// +// - "openai": cachedInput is a SUBSET of inTokens. The cached portion is +// billed at CachedInputPer1K (or InputPer1K when no override), and the +// non-cached remainder of inTokens at InputPer1K. cacheCreation is +// ignored (OpenAI has no analogue). +// - "anthropic", "bedrock": cachedInput (cache_read) and cacheCreation are +// ADDITIVE to inTokens. The three buckets are billed at CacheReadPer1K, +// CacheCreationPer1K, and InputPer1K respectively, each falling back +// to InputPer1K when the corresponding rate is zero. +// - Other surfaces: cached and cacheCreation are ignored; cost is +// inTokens*InputPer1K + outTokens*OutputPer1K. +func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, cacheCreation int64) Costs { // Clamp negatives to zero before any pricing math so a malformed // upstream count can never produce a negative cost. if inTokens < 0 { @@ -182,19 +195,8 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, if cacheCreation < 0 { cacheCreation = 0 } - if t == nil { - return Costs{}, false - } - byModel, ok := t.entries[provider] - if !ok { - return Costs{}, false - } - entry, ok := byModel[model] - if !ok { - return Costs{}, false - } output := (float64(outTokens) / 1000.0) * entry.OutputPer1K - switch provider { + switch surface { case "openai": // cachedInput is a subset of inTokens; clamp so a malformed // upstream (cached > total) can't produce a negative remainder. @@ -208,7 +210,7 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, } nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K cached := float64(clamped) / 1000.0 * cachedRate - return newCosts(nonCached, cached, 0, output), true + return newCosts(nonCached, cached, 0, output) case "anthropic", "bedrock": // Bedrock-Anthropic returns the same additive cache buckets as // first-party Anthropic; non-Anthropic Bedrock models simply report @@ -224,266 +226,9 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, input := float64(inTokens) / 1000.0 * entry.InputPer1K read := float64(cachedInput) / 1000.0 * readRate create := float64(cacheCreation) / 1000.0 * createRate - return newCosts(input, read, create, output), true + return newCosts(input, read, create, output) default: input := float64(inTokens) / 1000.0 * entry.InputPer1K - return newCosts(input, 0, 0, output), true + return newCosts(input, 0, 0, output) } } - -// Has reports whether the provider/model pair is present in the table. -func (t *Table) Has(provider, model string) bool { - if t == nil { - return false - } - byModel, ok := t.entries[provider] - if !ok { - return false - } - _, ok = byModel[model] - return ok -} - -// pricingFile mirrors the on-disk YAML schema. Keys are provider names; the -// nested map keys are model names. -type pricingFile map[string]map[string]struct { - InputPer1K float64 `yaml:"input_per_1k"` - OutputPer1K float64 `yaml:"output_per_1k"` - CachedInputPer1K float64 `yaml:"cached_input_per_1k"` - CacheReadPer1K float64 `yaml:"cache_read_per_1k"` - CacheCreationPer1K float64 `yaml:"cache_creation_per_1k"` -} - -const ( - // ReloadInterval is the mtime-poll cadence for the background reloader. - ReloadInterval = 30 * time.Second - - // errorBackoff bounds how often the loader logs a repeated parse error. - errorBackoff = 5 * time.Minute -) - -var basenameRegex = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`) - -// Loader is a confined, hot-reloadable pricing table reader. Construction -// must succeed against the target file; subsequent reload failures keep the -// previously-loaded table so callers never observe a blank price list. -type Loader struct { - baseDir string - fullPath string - pluginID string - table atomic.Pointer[Table] - mtime atomic.Int64 - failures metric.Int64Counter - interval time.Duration -} - -// NewLoader returns a pricing loader that overlays an optional file-based -// table on top of the embedded defaults. Missing override file, baseDir, or -// relPath is not an error: the loader keeps the embedded defaults so cost -// metadata is still emitted for known models. -// -// Errors: -// - bad basename, traversal segment, or absolute relPath are rejected so a -// misconfigured target surfaces immediately. -// - permission errors and YAML parse errors keep the defaults but log a -// warning; cost annotation does not silently break. -// -// failures is optional; pass nil in tests that do not care about -// reload-failure telemetry. -func NewLoader(baseDir, relPath, pluginID string, failures metric.Int64Counter) (*Loader, error) { - defaults := DefaultTable() - l := &Loader{ - baseDir: baseDir, - pluginID: pluginID, - failures: failures, - } - l.table.Store(defaults) - - if strings.TrimSpace(baseDir) == "" || strings.TrimSpace(relPath) == "" { - return l, nil - } - - full, err := resolveMiddlewareDataPath(baseDir, relPath) - if err != nil { - return nil, err - } - l.fullPath = full - - overlay, mtime, err := loadPricing(full) - if err != nil { - if errors.Is(err, fs.ErrNotExist) { - // Override file is optional. Defaults already stored. - return l, nil - } - // Symlink rejection, oversize file, parse failure, permission errors - // — surface so a misconfigured operator sees the problem instead of - // silently running with stale defaults. - return nil, fmt.Errorf("load pricing %s: %w", full, err) - } - l.table.Store(mergeOver(defaults, overlay)) - l.mtime.Store(mtime.UnixNano()) - return l, nil -} - -// Get returns the current pricing table. The returned pointer is immutable; -// callers must not mutate its contents. -func (l *Loader) Get() *Table { - if l == nil { - return nil - } - return l.table.Load() -} - -// WatchesFile reports whether this loader is bound to an override file on -// disk. False for defaults-only loaders (no operator override given). -// Callers use this to decide whether to spawn the mtime-poll goroutine. -func (l *Loader) WatchesFile() bool { - if l == nil { - return false - } - return l.fullPath != "" -} - -// SetReloadInterval overrides the mtime-poll cadence used by Reload. Calls -// after Reload has started have no effect on the running loop. Intended for -// tests; production code uses the default ReloadInterval. -func (l *Loader) SetReloadInterval(d time.Duration) { - if l == nil || d <= 0 { - return - } - l.interval = d -} - -// Reload runs a polling loop that checks the pricing file mtime every -// ReloadInterval (or the value passed to SetReloadInterval). Returns when -// ctx is cancelled. -func (l *Loader) Reload(ctx context.Context) { - if l == nil { - return - } - interval := l.interval - if interval <= 0 { - interval = ReloadInterval - } - t := time.NewTicker(interval) - defer t.Stop() - - var lastErrAt time.Time - for { - select { - case <-ctx.Done(): - return - case <-t.C: - if err := l.reload(); err != nil { - if l.failures != nil { - l.failures.Add(ctx, 1, metric.WithAttributes( - attribute.String("plugin", l.pluginID), - )) - } - now := time.Now() - if now.Sub(lastErrAt) >= errorBackoff { - log.Warnf("llmobs: pricing reload failed for %s: %v", l.fullPath, err) - lastErrAt = now - } - } - } - } -} - -// reload performs a single-shot mtime check and reload. The reloaded -// override file is merged on top of the embedded defaults; missing override -// (e.g. operator deleted the file) is not an error and reverts to defaults. -func (l *Loader) reload() error { - if l.fullPath == "" { - // Defaults-only loader; nothing on disk to reload. - return nil - } - mtime, err := statMtime(l.fullPath) - if err != nil { - if errors.Is(err, fs.ErrNotExist) { - // File was removed since startup. Drop back to defaults and - // reset mtime so a future re-creation triggers a reload. - l.table.Store(DefaultTable()) - l.mtime.Store(0) - return nil - } - return err - } - if mtime.UnixNano() == l.mtime.Load() { - return nil - } - - overlay, newMtime, err := loadPricing(l.fullPath) - if err != nil { - return err - } - l.table.Store(mergeOver(DefaultTable(), overlay)) - l.mtime.Store(newMtime.UnixNano()) - return nil -} - -// resolveMiddlewareDataPath validates relPath is a safe basename and resolves -// it under baseDir. An additional cleaned-prefix check guards against -// CVE-style edge cases where Join is used with trailing path segments. -func resolveMiddlewareDataPath(baseDir, relPath string) (string, error) { - if strings.TrimSpace(baseDir) == "" { - return "", errors.New("middleware-data-dir is not configured") - } - if relPath == "" { - return "", errors.New("pricing path is empty") - } - if !basenameRegex.MatchString(relPath) { - return "", fmt.Errorf("pricing path %q is not a safe basename", relPath) - } - if filepath.IsAbs(relPath) { - return "", fmt.Errorf("pricing path %q must be a basename, not absolute", relPath) - } - - cleanBase, err := filepath.Abs(filepath.Clean(baseDir)) - if err != nil { - return "", fmt.Errorf("resolve middleware-data-dir: %w", err) - } - full := filepath.Join(cleanBase, relPath) - cleanedFull := filepath.Clean(full) - if !strings.HasPrefix(cleanedFull, cleanBase+string(filepath.Separator)) && cleanedFull != cleanBase { - return "", fmt.Errorf("pricing path %q escapes middleware-data-dir", relPath) - } - return cleanedFull, nil -} - -func parsePricingBytes(data []byte) (*Table, error) { - dec := yaml.NewDecoder(bytes.NewReader(data)) - dec.KnownFields(true) - - var raw pricingFile - if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) { - return nil, fmt.Errorf("decode pricing yaml: %w", err) - } - - out := make(map[string]map[string]Entry, len(raw)) - for provider, models := range raw { - inner := make(map[string]Entry, len(models)) - for model, entry := range models { - for field, v := range map[string]float64{ - "input_per_1k": entry.InputPer1K, - "output_per_1k": entry.OutputPer1K, - "cached_input_per_1k": entry.CachedInputPer1K, - "cache_read_per_1k": entry.CacheReadPer1K, - "cache_creation_per_1k": entry.CacheCreationPer1K, - } { - if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) { - return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", provider, model, field, v) - } - } - inner[model] = Entry{ - InputPer1K: entry.InputPer1K, - OutputPer1K: entry.OutputPer1K, - CachedInputPer1K: entry.CachedInputPer1K, - CacheReadPer1K: entry.CacheReadPer1K, - CacheCreationPer1K: entry.CacheCreationPer1K, - } - } - out[provider] = inner - } - return &Table{entries: out}, nil -} diff --git a/proxy/internal/llm/pricing/pricing_other.go b/proxy/internal/llm/pricing/pricing_other.go deleted file mode 100644 index e65fffff1..000000000 --- a/proxy/internal/llm/pricing/pricing_other.go +++ /dev/null @@ -1,20 +0,0 @@ -//go:build !unix - -package pricing - -import ( - "fmt" - "time" -) - -// loadPricing is unavailable on non-Unix platforms because O_NOFOLLOW and -// fstat-from-FD are required to honour the spec's symlink-safety rules. The -// proxy is only deployed on Linux today; a Windows port would need an -// equivalent path-as-handle implementation. -func loadPricing(path string) (*Table, time.Time, error) { - return nil, time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path) -} - -func statMtime(path string) (time.Time, error) { - return time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path) -} diff --git a/proxy/internal/llm/pricing/pricing_test.go b/proxy/internal/llm/pricing/pricing_test.go index 7ac2a85dc..b946faa7f 100644 --- a/proxy/internal/llm/pricing/pricing_test.go +++ b/proxy/internal/llm/pricing/pricing_test.go @@ -1,47 +1,13 @@ -//go:build unix - package pricing import ( - "context" - "os" - "path/filepath" + "math" "testing" - "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func copyFixture(t *testing.T, src, dst string) { - t.Helper() - data, err := os.ReadFile(src) - require.NoError(t, err, "read source fixture") - require.NoError(t, os.WriteFile(dst, data, 0o600), "write target fixture") -} - -func TestNewLoader_HappyPath(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err, "NewLoader must succeed with a valid fixture") - table := l.Get() - require.NotNil(t, table, "table populated after load") - - cost, ok := table.Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "known provider/model resolves") - assert.InDelta(t, 0.00075, cost, 1e-9, "cost = 0.00015 + 0.0006 per 1k tokens") - - cost, ok = table.Cost("openai", "gpt-4o", 2000, 1000, 0, 0) - require.True(t, ok, "second known model resolves") - assert.InDelta(t, 0.015, cost, 1e-9, "cost for gpt-4o: 2*0.0025 + 1*0.01") - - cost, ok = table.Cost("anthropic", "claude-sonnet-4-5", 1000, 1000, 0, 0) - require.True(t, ok, "anthropic model resolves") - assert.InDelta(t, 0.018, cost, 1e-9, "cost for claude-sonnet-4-5: 0.003 + 0.015") -} - // TestCost_OpenAICachedSubsetDiscount proves OpenAI's cached input // tokens are billed at the configured cached_input_per_1k rate while // the non-cached remainder of input_tokens is billed at the regular @@ -65,11 +31,9 @@ func TestCost_OpenAICachedSubsetDiscount(t *testing.T) { "cached subset must bill at the discount rate; non-cached remainder at regular rate") } -// TestCost_OpenAICachedFallsBackToInputRate covers the operator -// opt-in contract: when CachedInputPer1K is unset (zero), cached -// tokens bill at the regular input rate. This matches today's -// behaviour (cached counts weren't extracted at all so they -// implicitly billed at the input rate via prompt_tokens). +// TestCost_OpenAICachedFallsBackToInputRate covers the fallback +// contract: when CachedInputPer1K is unset (zero), cached tokens bill +// at the regular input rate. func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) { tbl := &Table{entries: map[string]map[string]Entry{ "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}}, @@ -78,7 +42,7 @@ func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) { require.True(t, ok) want := 0.0025 + (500.0/1000.0)*0.01 assert.InDelta(t, want, cost, 1e-12, - "absent cached_input_per_1k rate must fall back to input_per_1k — same as pre-feature behaviour") + "absent cached_input_per_1k rate must fall back to input_per_1k") } // TestCost_OpenAIClampsCachedToInputCount is the defensive guard @@ -100,10 +64,7 @@ func TestCost_OpenAIClampsCachedToInputCount(t *testing.T) { // TestCost_AnthropicCacheReadAndCreationAreAdditive proves the // Anthropic shape: cache_read and cache_creation tokens are // ADDITIVE to input_tokens (not subset), each billed at its own -// configured rate. The two rates pull in opposite directions — -// cache_read is the cheaper read-from-cache rate (≈0.1× input), -// cache_creation is the more expensive write-to-cache rate -// (≈1.25× input). +// configured rate. func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) { tbl := &Table{entries: map[string]map[string]Entry{ "anthropic": {"claude-sonnet": { @@ -125,11 +86,9 @@ func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) { "each Anthropic input bucket must bill at its own configured rate") } -// TestCost_AnthropicCacheRatesFallBackToInput covers the no-opt-in +// TestCost_AnthropicCacheRatesFallBackToInput covers the no-rate // path: when neither CacheReadPer1K nor CacheCreationPer1K is set, -// cache tokens bill at the regular input rate. This is more -// accurate than today's behaviour (cache tokens ignored entirely) -// without requiring operators to opt in via YAML. +// cache tokens bill at the regular input rate. func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) { tbl := &Table{entries: map[string]map[string]Entry{ "anthropic": {"claude-sonnet": {InputPer1K: 0.003, OutputPer1K: 0.015}}, @@ -139,259 +98,39 @@ func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) { // Without overrides: every input bucket at input_per_1k. want := ((256.0+768.0+512.0)/1000.0)*0.003 + (200.0/1000.0)*0.015 assert.InDelta(t, want, cost, 1e-12, - "absent cache rates must fall back to input_per_1k — Anthropic cache tokens were ignored before this change, billing at input rate is more accurate as a default") + "absent cache rates must fall back to input_per_1k") } -func TestNewLoader_UnknownModel(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) +// TestEntryCosts_SurfaceSelectsFormula pins that the formula branches on +// the SURFACE, not on which table the entry came from: the same entry +// bills a subset carve-out on "openai", additive buckets on +// "anthropic"/"bedrock", and ignores cache counts everywhere else. This +// is what keeps per-provider-record entries (looked up by record id) +// mathematically identical to defaults-table entries. +func TestEntryCosts_SurfaceSelectsFormula(t *testing.T) { + e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01, CachedInputPer1K: 0.001, CacheReadPer1K: 0.0002, CacheCreationPer1K: 0.0025} - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) + openai := EntryCosts(e, "openai", 1000, 0, 400, 300) + assert.InDelta(t, (600.0/1000.0)*0.002+(400.0/1000.0)*0.001, openai.TotalUSD, 1e-12, + "openai: cached is a subset, cacheCreation ignored") - _, ok := l.Get().Cost("openai", "fantasy-model", 10, 10, 0, 0) - assert.False(t, ok, "unknown model returns ok=false") + anthropic := EntryCosts(e, "anthropic", 1000, 0, 400, 300) + assert.InDelta(t, 0.002+(400.0/1000.0)*0.0002+(300.0/1000.0)*0.0025, anthropic.TotalUSD, 1e-12, + "anthropic: cache buckets are additive") - _, ok = l.Get().Cost("cohere", "anything", 10, 10, 0, 0) - assert.False(t, ok, "unknown provider returns ok=false") + bedrock := EntryCosts(e, "bedrock", 1000, 0, 400, 300) + assert.InDelta(t, anthropic.TotalUSD, bedrock.TotalUSD, 1e-12, "bedrock shares the anthropic formula") + + other := EntryCosts(e, "gemini", 1000, 0, 400, 300) + assert.InDelta(t, 0.002, other.TotalUSD, 1e-12, "unknown surface: cache counts ignored") } -func TestNewLoader_InvalidYAMLRejected(t *testing.T) { - base := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(base, "pricing.yaml"), []byte("\t- this is not: valid: yaml: :["), 0o600)) - - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "invalid YAML must surface as construction error") -} - -func TestLoader_ReloadKeepsPreviousOnParseError(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - require.NotNil(t, l.Get(), "initial table populated") - - // Overwrite with content that violates the strict schema (extra field) - // plus a bumped mtime to trigger reload. - require.NoError(t, os.WriteFile(target, []byte("openai:\n gpt-4o:\n input_per_1k: 1.0\n output_per_1k: 2.0\n bogus_field: nope\n"), 0o600)) - future := time.Now().Add(time.Hour) - require.NoError(t, os.Chtimes(target, future, future)) - - err = l.reload() - require.Error(t, err, "parse error surfaced by reload()") - - cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "previous table still available after parse failure") - assert.InDelta(t, 0.00075, cost, 1e-9, "previous cost preserved") -} - -func TestLoader_ReloadNoChangeIsNoOp(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - ptrBefore := l.Get() - - require.NoError(t, l.reload(), "no-change reload must not error") - ptrAfter := l.Get() - assert.Same(t, ptrBefore, ptrAfter, "table pointer unchanged when mtime unchanged") -} - -func TestLoader_ReloadDetectsChange(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - - updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n") - require.NoError(t, os.WriteFile(target, updated, 0o600)) - future := time.Now().Add(time.Hour) - require.NoError(t, os.Chtimes(target, future, future)) - - require.NoError(t, l.reload(), "reload must succeed on valid new content") - - cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "updated model still present") - assert.InDelta(t, 3.0, cost, 0.0001, "new prices are applied: 1 + 2 per 1k") -} - -// TestLoader_ReloadGoroutinePicksUpChanges proves the background goroutine -// started via Reload actually swaps the pricing table when the file changes -// on disk. Without that goroutine running, pricing edits would never reach -// requests until a proxy restart. -func TestLoader_ReloadGoroutinePicksUpChanges(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - l.SetReloadInterval(20 * time.Millisecond) - - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - done := make(chan struct{}) - go func() { - l.Reload(ctx) - close(done) - }() - - // Before any rewrite, the loader holds the fixture's prices. - costBefore, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - require.True(t, ok, "fixture model must resolve initially") - assert.InDelta(t, 0.00075, costBefore, 1e-9, "fixture prices apply before rewrite") - - updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n") - require.NoError(t, os.WriteFile(target, updated, 0o600)) - future := time.Now().Add(time.Hour) - require.NoError(t, os.Chtimes(target, future, future)) - - deadline := time.Now().Add(2 * time.Second) - for { - cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - if ok && cost > 2.5 { - break - } - if time.Now().After(deadline) { - t.Fatalf("background reloader did not pick up rewrite within deadline") - } - time.Sleep(10 * time.Millisecond) - } - - cancel() - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("Reload loop did not exit after cancel") - } -} - -func TestLoader_ReloadBackgroundLoopCancellation(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.NoError(t, err) - - ctx, cancel := context.WithCancel(context.Background()) - done := make(chan struct{}) - go func() { - l.Reload(ctx) - close(done) - }() - cancel() - - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("Reload loop did not exit on context cancel") - } -} - -func TestNewLoader_PathValidation(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - - cases := []struct { - name string - relPath string - }{ - {"traversal", "../../etc/passwd"}, - {"absolute", "/etc/passwd"}, - {"slash in basename", "sub/pricing.yaml"}, - {"control chars", "pricing\x00.yaml"}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - _, err := NewLoader(base, tc.relPath, "llm_observability", nil) - require.Error(t, err, "NewLoader must reject %q", tc.relPath) - }) - } - - // Empty relPath is no longer a validation error: the loader treats it - // as "no override file, defaults only" so cost metadata is still - // emitted for the embedded models out of the box. - t.Run("empty falls back to defaults", func(t *testing.T) { - l, err := NewLoader(base, "", "llm_observability", nil) - require.NoError(t, err, "empty relPath should yield a defaults-only loader") - require.NotNil(t, l, "loader must be returned") - require.False(t, l.WatchesFile(), "no file watching when no override is given") - _, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0) - assert.True(t, ok, "embedded defaults should still resolve gpt-4o-mini") - }) -} - -// TestNewLoader_PathValidation_Extended covers the remaining attack shapes -// called out in C2: dot references, embedded traversal segments, and a -// newline in the basename. The basename regex must reject each one even -// though filepath.Clean would otherwise collapse them. -func TestNewLoader_PathValidation_Extended(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml")) - - cases := []struct { - name string - relPath string - }{ - {"dot", "."}, - {"dotdot", ".."}, - {"relative traversal", "../pricing.yaml"}, - {"embedded slash", "pri/cing.yaml"}, - {"newline", "pricing\n.yaml"}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - _, err := NewLoader(base, tc.relPath, "llm_observability", nil) - require.Error(t, err, "NewLoader must reject %q", tc.relPath) - }) - } -} - -// TestNewLoader_ValidBasenameLoads proves the allowlist is exclusive: a -// basename containing only safe characters under baseDir loads. Without this -// a regression that over-tightened the regex would silently break valid -// deployments. -func TestNewLoader_ValidBasenameLoads(t *testing.T) { - base := t.TempDir() - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing-v2_prod.yaml")) - - l, err := NewLoader(base, "pricing-v2_prod.yaml", "llm_observability", nil) - require.NoError(t, err, "basename with _, -, . must load") - require.NotNil(t, l.Get(), "table populated") -} - -// TestNewLoader_SymlinkOutsideBaseDirRejected constructs a symlink under -// baseDir that points to a file outside it. O_NOFOLLOW must refuse to open -// the symlink even though the symlink path itself is a valid basename under -// baseDir. -func TestNewLoader_SymlinkOutsideBaseDirRejected(t *testing.T) { - outside := t.TempDir() - target := filepath.Join(outside, "evil.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target) - - base := t.TempDir() - link := filepath.Join(base, "pricing.yaml") - require.NoError(t, os.Symlink(target, link), "symlink setup") - - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "O_NOFOLLOW must reject symlink even when it points outside baseDir") -} - -func TestNewLoader_SymlinkRejected(t *testing.T) { - base := t.TempDir() - concrete := filepath.Join(base, "real.yaml") - copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), concrete) - - link := filepath.Join(base, "pricing.yaml") - require.NoError(t, os.Symlink(concrete, link), "symlink setup") - - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "O_NOFOLLOW must reject symlinked targets") +// TestEntryCosts_ClampsNegativeTokens: malformed upstream counts must +// never produce a negative cost. +func TestEntryCosts_ClampsNegativeTokens(t *testing.T) { + e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01} + c := EntryCosts(e, "openai", -50, -10, -5, -3) + assert.Zero(t, c.TotalUSD, "all-negative counts clamp to zero cost") } func TestTableCost_NilSafe(t *testing.T) { @@ -402,31 +141,37 @@ func TestTableCost_NilSafe(t *testing.T) { assert.False(t, t1.Has("x", "y"), "nil table has nothing") } -func TestLoaderGet_NilSafe(t *testing.T) { - var l *Loader - assert.Nil(t, l.Get(), "nil loader returns nil table") -} - -// TestNewLoader_RejectsOversizedFile_FixesM4 proves the loader bounds reads -// at maxPricingBytes so a hostile file cannot exhaust process memory. -func TestNewLoader_RejectsOversizedFile_FixesM4(t *testing.T) { - base := t.TempDir() - target := filepath.Join(base, "pricing.yaml") - - // Build a YAML payload larger than the cap. We pad with valid YAML - // comments so a partial read would still fail the size check rather - // than the parser. - header := "openai:\n" - bigComment := make([]byte, maxPricingBytes+1024) - for i := range bigComment { - bigComment[i] = ' ' +func TestNewTable_ValidatesRates(t *testing.T) { + good := map[string]map[string]EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125}}, } - bigComment[0] = '#' - bigComment[len(bigComment)-1] = '\n' - payload := append([]byte(header), bigComment...) - require.NoError(t, os.WriteFile(target, payload, 0o600)) + tbl, err := NewTable(good) + require.NoError(t, err) + cost, ok := tbl.Cost("openai", "gpt-4o", 1000, 1000, 0, 0) + require.True(t, ok, "entry survives the wire conversion") + assert.InDelta(t, 0.0125, cost, 1e-9) - _, err := NewLoader(base, "pricing.yaml", "llm_observability", nil) - require.Error(t, err, "oversized pricing file must be rejected") - assert.Contains(t, err.Error(), "exceeds", "rejection must reference the byte cap") + _, ok = tbl.Cost("openai", "unknown-model", 1, 1, 0, 0) + assert.False(t, ok, "unknown model misses") + + for name, bad := range map[string]EntryJSON{ + "negative input": {InputPer1K: -1, OutputPer1K: 0.01}, + "NaN output": {InputPer1K: 0.01, OutputPer1K: math.NaN()}, + "Inf cache read": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: math.Inf(1)}, + "negative cached": {InputPer1K: 0.01, OutputPer1K: 0.01, CachedInputPer1K: -0.001}, + } { + _, err := NewTable(map[string]map[string]EntryJSON{"openai": {"m": bad}}) + assert.Error(t, err, "case %q must be rejected so a corrupt config fails the chain build instead of mispricing", name) + } +} + +func TestNewTable_NilAndEmpty(t *testing.T) { + tbl, err := NewTable(nil) + require.NoError(t, err, "nil map builds an empty (never-matching) table") + _, ok := tbl.Cost("openai", "gpt-4o", 1, 1, 0, 0) + assert.False(t, ok, "empty table prices nothing") + + entries, err := NewEntries(nil) + require.NoError(t, err) + assert.Empty(t, entries, "nil in, empty (never-matching) map out for the per-record map") } diff --git a/proxy/internal/llm/pricing/pricing_unix.go b/proxy/internal/llm/pricing/pricing_unix.go deleted file mode 100644 index 4f3ea33a2..000000000 --- a/proxy/internal/llm/pricing/pricing_unix.go +++ /dev/null @@ -1,68 +0,0 @@ -//go:build unix - -package pricing - -import ( - "fmt" - "io" - "os" - "syscall" - "time" - - log "github.com/sirupsen/logrus" -) - -// maxPricingBytes caps the size of the pricing YAML on read so a hostile or -// runaway file cannot exhaust process memory during reload. 1 MiB is several -// orders of magnitude larger than any reasonable pricing table. -const maxPricingBytes int64 = 1 << 20 - -// loadPricing opens the file with O_NOFOLLOW, fstats the open descriptor, -// and parses from that same descriptor. Never re-opens by path so a -// mid-read rename or symlink swap cannot substitute content. Bytes are -// capped at maxPricingBytes so the loader cannot be coerced into reading an -// unbounded file. -func loadPricing(path string) (*Table, time.Time, error) { - f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0) - if err != nil { - return nil, time.Time{}, fmt.Errorf("open %s: %w", path, err) - } - defer func() { - if cerr := f.Close(); cerr != nil { - log.Debugf("close pricing file %s: %v", path, cerr) - } - }() - - info, err := f.Stat() - if err != nil { - return nil, time.Time{}, fmt.Errorf("fstat %s: %w", path, err) - } - if !info.Mode().IsRegular() { - return nil, time.Time{}, fmt.Errorf("pricing file %s is not a regular file", path) - } - - data, err := io.ReadAll(io.LimitReader(f, maxPricingBytes+1)) - if err != nil { - return nil, time.Time{}, fmt.Errorf("read %s: %w", path, err) - } - if int64(len(data)) > maxPricingBytes { - return nil, time.Time{}, fmt.Errorf("pricing file %s exceeds %d bytes", path, maxPricingBytes) - } - - table, err := parsePricingBytes(data) - if err != nil { - return nil, time.Time{}, err - } - return table, info.ModTime(), nil -} - -// statMtime returns the mtime of the file at path. It uses lstat semantics -// via os.Lstat so a symlink swap is detected even though O_NOFOLLOW will -// later reject the open. -func statMtime(path string) (time.Time, error) { - info, err := os.Lstat(path) - if err != nil { - return time.Time{}, fmt.Errorf("lstat %s: %w", path, err) - } - return info.ModTime(), nil -} diff --git a/proxy/internal/middleware/builtin/builtin.go b/proxy/internal/middleware/builtin/builtin.go index 9ea4cf89d..9df60dd65 100644 --- a/proxy/internal/middleware/builtin/builtin.go +++ b/proxy/internal/middleware/builtin/builtin.go @@ -36,15 +36,13 @@ var defaultRegistry = middleware.NewRegistry() // FactoryContext is the per-process bag that concrete factories may // consult during construction. It carries the proxy-lifetime context, -// the data directory used for static config files (pricing tables, -// allowlists), the OTel meter, and the proxy logger. +// the OTel meter, and the proxy logger. // // Configure must be called once at boot before any chain build calls // Resolve. Calling it twice overwrites the prior value; tests may rely // on this to reset state. type FactoryContext struct { Context context.Context - DataDir string Meter metric.Meter Logger *log.Logger MgmtClient MgmtClient @@ -58,12 +56,11 @@ var ( // Configure stores the per-process FactoryContext. Concrete factories // reach for it via Context(). mgmt may be nil on tests / standalone // builds with no management server; consumers must guard. -func Configure(ctx context.Context, dataDir string, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) { +func Configure(ctx context.Context, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) { ctxMu.Lock() defer ctxMu.Unlock() ctxStore = FactoryContext{ Context: ctx, - DataDir: dataDir, Meter: meter, Logger: logger, MgmtClient: mgmt, diff --git a/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go b/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go index f479eb563..be3682c0e 100644 --- a/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go +++ b/proxy/internal/middleware/builtin/cost_calculation_matrix_test.go @@ -12,6 +12,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + mgmtpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing" "github.com/netbirdio/netbird/proxy/internal/middleware" "github.com/netbirdio/netbird/proxy/internal/middleware/builtin" "github.com/netbirdio/netbird/proxy/internal/middleware/builtin/cost_meter" @@ -19,17 +20,20 @@ import ( "github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_response_parser" ) -// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the embedded default pricing -// table and asserts exact USD amounts hardcoded from the vendors' published prices, including the cache split. +// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the REAL default pricing +// table management ships (mgmtpricing.DefaultTable, catalog-derived) and asserts exact USD amounts hardcoded from +// the vendors' published prices, including the cache split. This is the cross-stack pricing contract test: the +// management-side Entry JSON must decode into the proxy-side table and produce these exact costs. func TestCostCalculation_ProviderMatrix(t *testing.T) { - // Empty data dir → embedded defaults, like a proxy with no pricing override. - builtin.Configure(context.Background(), t.TempDir(), nil, nil, nil) + builtin.Configure(context.Background(), nil, nil, nil) reqMW, err := llm_request_parser.Factory{}.New(nil) require.NoError(t, err, "build llm_request_parser") respMW, err := llm_response_parser.Factory{}.New(nil) require.NoError(t, err, "build llm_response_parser") - costMW, err := cost_meter.Factory{}.New(nil) + costCfgJSON, err := json.Marshal(map[string]any{"pricing": map[string]any{"defaults": mgmtpricing.DefaultTable()}}) + require.NoError(t, err, "marshal management default table into cost_meter config") + costMW, err := cost_meter.Factory{}.New(costCfgJSON) require.NoError(t, err, "build cost_meter") t.Cleanup(func() { _ = costMW.Close() }) diff --git a/proxy/internal/middleware/builtin/cost_meter/factory.go b/proxy/internal/middleware/builtin/cost_meter/factory.go index b8a58d10e..2993ce32e 100644 --- a/proxy/internal/middleware/builtin/cost_meter/factory.go +++ b/proxy/internal/middleware/builtin/cost_meter/factory.go @@ -2,7 +2,6 @@ package cost_meter import ( "bytes" - "context" "encoding/json" "fmt" @@ -11,16 +10,27 @@ import ( "github.com/netbirdio/netbird/proxy/internal/middleware/builtin" ) -// defaultPricingFilename is the basename probed inside the proxy data -// directory when no override is configured. -const defaultPricingFilename = "pricing.yaml" - -// Config is the on-wire configuration for the middleware. +// Config is the on-wire configuration for the middleware, synthesized by +// management (buildCostMeterConfigJSON). The proxy has no embedded price +// list: this payload is the only pricing source, and updates arrive as +// ordinary mapping pushes that rebuild the chain (and with it this +// middleware instance) — no per-request fetches, no reload loops. type Config struct { - // PricingPath optionally overrides the basename of the pricing - // file probed inside the proxy data directory. When empty the - // loader falls back to "pricing.yaml". - PricingPath string `json:"pricing_path"` + Pricing *PricingConfig `json:"pricing"` +} + +// PricingConfig carries the full pricing table: +// - Defaults: parser surface ("openai"/"anthropic"/"bedrock") -> +// normalized model id -> rates, matched against llm.provider + +// llm.model. +// - Providers: provider record id -> normalized model id -> rates, +// matched against the llm.resolved_provider_id metadata llm_router +// stamps. Entries arrive fully materialized (management folds default +// cache rates in at synth time), so lookup order is simply +// per-record first, defaults second. +type PricingConfig struct { + Defaults map[string]map[string]pricing.EntryJSON `json:"defaults"` + Providers map[string]map[string]pricing.EntryJSON `json:"providers"` } // Factory builds cost_meter instances from raw config bytes. @@ -29,45 +39,45 @@ type Factory struct{} // ID returns the registry identifier. func (Factory) ID() string { return ID } -// New constructs a middleware instance. Empty, null, and {} configs -// are accepted; non-empty rawConfig that fails to unmarshal is -// rejected so misconfigurations surface at chain build time. The -// pricing loader is built once per instance and reused across -// invocations. +// New constructs a middleware instance. Empty, null, and {} configs are +// accepted for backward compatibility with a management server that +// predates config-delivered pricing — the instance then skips every cost +// computation (unknown_model) and a warning is logged once at build time. +// Non-empty rawConfig that fails to unmarshal, or a table carrying a +// non-finite / negative rate, is rejected so misconfigurations surface at +// chain build time. func (Factory) New(rawConfig []byte) (middleware.Middleware, error) { cfg, err := decodeConfig(rawConfig) if err != nil { return nil, err } - fctx := builtin.Context() - pricingPath := cfg.PricingPath - if pricingPath == "" { - pricingPath = defaultPricingFilename + if cfg.Pricing == nil { + if logger := builtin.Context().Logger; logger != nil { + logger.Warnf("cost_meter: no pricing table in middleware config; management predates config-delivered pricing — every request will record cost.skipped=unknown_model ($0)") + } + return newMiddleware(mustEmptyTable(), nil), nil } - loader, err := pricing.NewLoader(fctx.DataDir, pricingPath, ID, nil) + defaults, err := pricing.NewTable(cfg.Pricing.Defaults) if err != nil { - return nil, fmt.Errorf("init pricing loader: %w", err) + return nil, fmt.Errorf("cost_meter pricing defaults: %w", err) } - - cancel := startReloader(fctx.Context, loader) - - return newMiddleware(loader, cancel), nil + perRecord, err := pricing.NewEntries(cfg.Pricing.Providers) + if err != nil { + return nil, fmt.Errorf("cost_meter per-provider pricing: %w", err) + } + return newMiddleware(defaults, perRecord), nil } -// startReloader binds the loader's mtime-poll goroutine to a context -// derived from the proxy-lifetime context and returns its cancel func so -// the owning middleware can stop the goroutine on teardown. Returns nil -// when there's nothing to watch (nil context or defaults-only loader), in -// which case the middleware's Close is a no-op. -func startReloader(ctx context.Context, loader *pricing.Loader) context.CancelFunc { - if ctx == nil || !loader.WatchesFile() { - return nil +// mustEmptyTable returns a valid empty table. NewTable on a nil map cannot +// fail; the panic guard documents that invariant. +func mustEmptyTable() *pricing.Table { + t, err := pricing.NewTable(nil) + if err != nil { + panic(fmt.Sprintf("cost_meter: empty pricing table must build: %v", err)) } - cctx, cancel := context.WithCancel(ctx) - go loader.Reload(cctx) - return cancel + return t } // decodeConfig accepts empty, null, and {} configs, returning a diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware.go b/proxy/internal/middleware/builtin/cost_meter/middleware.go index 63da6d17b..2ce706cda 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware.go @@ -1,7 +1,9 @@ // Package cost_meter implements the SlotOnResponse middleware that // converts token-usage metadata emitted by llm_response_parser into a -// per-request USD cost estimate. The middleware uses the shared pricing -// loader so operator pricing overrides apply to the chain. +// per-request USD cost estimate. Pricing arrives from management inside +// the middleware config: a per-provider-record table (the operator's +// stored prices, matched via llm.resolved_provider_id) consulted first, +// then the surface-keyed defaults table. package cost_meter import ( @@ -17,7 +19,9 @@ import ( const ID = "cost_meter" // Version is the implementation version emitted via the spec merge. -const Version = "1.0.0" +// 1.1.0: pricing is config-delivered (defaults + per-provider-record +// entries) instead of proxy-embedded. +const Version = "1.1.0" // Skip reasons emitted under KeyCostSkipped. The set is closed; the // dashboard surfaces these verbatim. @@ -42,19 +46,21 @@ var metadataKeys = []string{ } // Middleware computes a per-response cost estimate from the token -// counts emitted upstream by llm_response_parser. +// counts emitted upstream by llm_response_parser. Both tables are +// immutable — a pricing change arrives as a mapping push that rebuilds +// the chain with a fresh instance. type Middleware struct { - loader *pricing.Loader - // cancel stops this instance's pricing-reload goroutine. Non-nil only - // when the loader watches an override file; Close calls it so a chain - // rebuild doesn't leak a poll goroutine per retired instance. - cancel context.CancelFunc + // defaults is the surface-keyed table (llm.provider x llm.model). + defaults *pricing.Table + // perRecord is keyed by provider record id (llm.resolved_provider_id) + // then normalized model id; entries arrive fully materialized from + // management. Consulted before defaults. May be nil. + perRecord map[string]map[string]pricing.Entry } -// newMiddleware constructs a Middleware bound to the given pricing loader. -// cancel may be nil (defaults-only loader with no reloader to stop). -func newMiddleware(loader *pricing.Loader, cancel context.CancelFunc) *Middleware { - return &Middleware{loader: loader, cancel: cancel} +// newMiddleware constructs a Middleware over the given pricing tables. +func newMiddleware(defaults *pricing.Table, perRecord map[string]map[string]pricing.Entry) *Middleware { + return &Middleware{defaults: defaults, perRecord: perRecord} } // ID returns the registry identifier. @@ -79,16 +85,9 @@ func (m *Middleware) MetadataKeys() []string { // response. func (m *Middleware) MutationsSupported() bool { return false } -// Close stops this instance's pricing-reload goroutine, if any. Called by -// the chain when a rebuild retires the instance, so the mtime-poll loop -// doesn't outlive the chain it belonged to. Safe to call on a nil receiver -// and on an instance with no reloader. -func (m *Middleware) Close() error { - if m != nil && m.cancel != nil { - m.cancel() - } - return nil -} +// Close releases resources owned by the middleware. Stateless — the +// pricing tables are plain maps owned by this instance. +func (m *Middleware) Close() error { return nil } // Invoke reads provider, model, and token metadata, looks up pricing, // and emits either KeyCostUSDTotal or KeyCostSkipped. The decision is @@ -144,8 +143,7 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar return out, nil } - table := m.loader.Get() - costs, ok := table.Costs(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens) + costs, ok := m.lookupCosts(in.Metadata, provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens) if !ok { out.Metadata = skip(skipUnknownModel) return out, nil @@ -164,6 +162,26 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar return out, nil } +// lookupCosts resolves the price for this request and computes the cost +// split. Resolution order: +// +// 1. Per-provider-record entry: the operator's stored price for the +// provider route that served the request, keyed by the +// llm.resolved_provider_id metadata llm_router stamped on the allow +// path. Absent metadata (e.g. no router in the chain) skips this tier. +// 2. Surface defaults: the catalog-derived table keyed by llm.provider. +// +// The surface always selects the cache formula — a per-record entry for an +// Anthropic route still bills its cache buckets additively. +func (m *Middleware) lookupCosts(md []middleware.KV, surface, model string, inTokens, outTokens, cachedTokens, cacheCreationTokens int64) (pricing.Costs, bool) { + if recordID := lookupKV(md, middleware.KeyLLMResolvedProviderID); recordID != "" { + if entry, ok := m.perRecord[recordID][model]; ok { + return pricing.EntryCosts(entry, surface, inTokens, outTokens, cachedTokens, cacheCreationTokens), true + } + } + return m.defaults.Costs(surface, model, inTokens, outTokens, cachedTokens, cacheCreationTokens) +} + // usd renders a cost as the fixed-precision string every cost.usd_* key // carries, so the per-bucket values and the aggregates round identically. // diff --git a/proxy/internal/middleware/builtin/cost_meter/middleware_test.go b/proxy/internal/middleware/builtin/cost_meter/middleware_test.go index e5d431d77..482061270 100644 --- a/proxy/internal/middleware/builtin/cost_meter/middleware_test.go +++ b/proxy/internal/middleware/builtin/cost_meter/middleware_test.go @@ -3,39 +3,50 @@ package cost_meter import ( "context" "encoding/json" - "os" - "path/filepath" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/netbirdio/netbird/proxy/internal/llm/pricing" "github.com/netbirdio/netbird/proxy/internal/middleware" - "github.com/netbirdio/netbird/proxy/internal/middleware/builtin" ) -const fixturePricing = `openai: - gpt-4o: - input_per_1k: 0.0025 - output_per_1k: 0.01 - gpt-4o-mini: - input_per_1k: 0.00015 - output_per_1k: 0.0006 -anthropic: - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 -` - -// configureBuiltin points the package-level FactoryContext at a tmp -// directory containing the test pricing fixture. Returns the path so -// callers can override files later if needed. -func configureBuiltin(t *testing.T) string { +// fixtureConfig mirrors what management's buildCostMeterConfigJSON ships: +// a surface-keyed defaults table. Rates match the retired YAML fixture so +// every cost assertion below is byte-identical to the pre-feature values. +func fixtureConfig(t *testing.T) []byte { t.Helper() - dir := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricing), 0o600), "write pricing fixture") - builtin.Configure(context.Background(), dir, nil, nil, nil) - return dir + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": { + "gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}, + "gpt-4o-mini": {InputPer1K: 0.00015, OutputPer1K: 0.0006}, + }, + "anthropic": { + "claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015}, + }, + }, + }}) + require.NoError(t, err) + return raw +} + +// fixtureConfigWithCache adds the cache-rate fields. +func fixtureConfigWithCache(t *testing.T) []byte { + t.Helper() + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": { + "gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125}, + }, + "anthropic": { + "claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375}, + }, + }, + }}) + require.NoError(t, err) + return raw } func metaValue(t *testing.T, kvs []middleware.KV, key string) (string, bool) { @@ -56,8 +67,7 @@ func buildMiddleware(t *testing.T, raw []byte) middleware.Middleware { } func TestMiddleware_StaticSurface(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) assert.Equal(t, ID, mw.ID(), "ID must match the registered constant") assert.Equal(t, Version, mw.Version(), "Version must match the constant") @@ -79,8 +89,10 @@ func TestMiddleware_StaticSurface(t *testing.T) { assert.Equal(t, expected, keys, "metadata key allowlist must match the spec") } +// TestFactory_AcceptsEmptyAndJSONConfig: empty/null/{} configs are what an +// old management (pre config-delivered pricing) sends — they must build a +// working (all-skip) instance, never fail the chain. func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) { - configureBuiltin(t) cases := [][]byte{nil, {}, []byte("null"), []byte("{}"), []byte(" ")} for _, raw := range cases { mw, err := Factory{}.New(raw) @@ -90,15 +102,57 @@ func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) { } func TestFactory_RejectsMalformedConfig(t *testing.T) { - configureBuiltin(t) mw, err := Factory{}.New([]byte("{not json")) require.Error(t, err, "malformed config must surface at construction") assert.Nil(t, mw, "no instance is returned on error") } -func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) +// TestFactory_RejectsInvalidRates: a non-finite or negative rate anywhere +// in the table fails the chain build (defense-in-depth behind management's +// API validation) rather than silently mispricing. +func TestFactory_RejectsInvalidRates(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: -0.0025, OutputPer1K: 0.01}}, + }, + }}) + require.NoError(t, err) + mw, err := Factory{}.New(raw) + require.Error(t, err, "negative rate must fail the build") + assert.Nil(t, mw) + + raw, err = json.Marshal(Config{Pricing: &PricingConfig{ + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-1": {"m": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: -1}}, + }, + }}) + require.NoError(t, err) + _, err = Factory{}.New(raw) + require.Error(t, err, "per-record tables validate too") +} + +// TestFactory_NilPricingSkipsEverything is the version-skew contract: a +// new proxy under an old management ({} config) must build, allow, and +// skip with unknown_model — degraded but never broken. +func TestFactory_NilPricingSkipsEverything(t *testing.T) { + mw := buildMiddleware(t, []byte("{}")) + out, err := mw.Invoke(context.Background(), &middleware.Input{ + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMProvider, Value: "openai"}, + {Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + {Key: middleware.KeyLLMInputTokens, Value: "1000"}, + {Key: middleware.KeyLLMOutputTokens, Value: "1000"}, + }, + }) + require.NoError(t, err) + assert.Equal(t, middleware.DecisionAllow, out.Decision, "cost_meter always allows") + value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped) + require.True(t, ok, "no pricing table means every request skips") + assert.Equal(t, skipUnknownModel, value) +} + +func TestFactory_ConfigDefaultsPriceRequests(t *testing.T) { + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -116,34 +170,78 @@ func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) { assert.Equal(t, "0.000750000", value, "0.00015 + 0.0006 per 1k tokens, 9-decimal format") } -func TestFactory_PricingPathOverride(t *testing.T) { - dir := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(dir, "custom.yaml"), []byte(fixturePricing), 0o600), "write custom pricing") - builtin.Configure(context.Background(), dir, nil, nil, nil) - - raw, err := json.Marshal(Config{PricingPath: "custom.yaml"}) +// TestInvoke_PerRecordEntryBeatsDefaults: when llm_router resolved a +// provider record whose operator pinned a price for the model, that price +// wins over the surface default. +func TestInvoke_PerRecordEntryBeatsDefaults(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}}, + }, + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-azure": {"gpt-4o": {InputPer1K: 0.005, OutputPer1K: 0.02}}, + }, + }}) require.NoError(t, err) - mw := buildMiddleware(t, raw) + out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ {Key: middleware.KeyLLMProvider, Value: "openai"}, {Key: middleware.KeyLLMModel, Value: "gpt-4o"}, - {Key: middleware.KeyLLMInputTokens, Value: "2000"}, + {Key: middleware.KeyLLMResolvedProviderID, Value: "prov-azure"}, + {Key: middleware.KeyLLMInputTokens, Value: "1000"}, {Key: middleware.KeyLLMOutputTokens, Value: "1000"}, }, }) require.NoError(t, err) - value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) - require.True(t, ok, "cost.usd_total must be emitted with custom pricing path") - assert.Equal(t, "0.015000000", value, "2*0.0025 + 1*0.01 = 0.015 with 9-decimal format") + require.True(t, ok) + assert.Equal(t, "0.025000000", value, "operator's per-record price (0.005+0.02) wins over the default (0.0025+0.01)") } -func TestInvoke_ComputesCostForKnownModel(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) +// TestInvoke_PerRecordMissFallsBackToDefaults: a resolved record with no +// entry for this model (or no entries at all) falls through to the +// surface defaults — gateway providers rely on exactly this. +func TestInvoke_PerRecordMissFallsBackToDefaults(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Defaults: map[string]map[string]pricing.EntryJSON{ + "openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}}, + }, + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-1": {"some-other-model": {InputPer1K: 1, OutputPer1K: 1}}, + }, + }}) + require.NoError(t, err) + mw := buildMiddleware(t, raw) + for name, recordID := range map[string]string{ + "record with other models": "prov-1", + "record with no entries": "prov-gateway", + } { + t.Run(name, func(t *testing.T) { + out, err := mw.Invoke(context.Background(), &middleware.Input{ + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMProvider, Value: "openai"}, + {Key: middleware.KeyLLMModel, Value: "gpt-4o"}, + {Key: middleware.KeyLLMResolvedProviderID, Value: recordID}, + {Key: middleware.KeyLLMInputTokens, Value: "1000"}, + {Key: middleware.KeyLLMOutputTokens, Value: "1000"}, + }, + }) + require.NoError(t, err) + value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) + require.True(t, ok, "per-record miss must fall back to the surface default, not skip") + assert.Equal(t, "0.012500000", value, "default rates apply") + }) + } +} + +// TestInvoke_NoResolvedProviderIDUsesDefaults: metadata without a +// resolved provider id (router denied, or a chain without llm_router) +// prices from the defaults table directly. +func TestInvoke_NoResolvedProviderIDUsesDefaults(t *testing.T) { + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ {Key: middleware.KeyLLMProvider, Value: "anthropic"}, @@ -153,17 +251,15 @@ func TestInvoke_ComputesCostForKnownModel(t *testing.T) { }, }) require.NoError(t, err) - value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) - require.True(t, ok, "cost.usd_total must be emitted") + require.True(t, ok) assert.Equal(t, "0.018000000", value, "0.003 + 0.015 = 0.018 with 9-decimal format") _, skipped := metaValue(t, out.Metadata, middleware.KeyCostSkipped) assert.False(t, skipped, "cost.skipped must not be set when cost is computed") } func TestInvoke_MissingProvider(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -179,8 +275,7 @@ func TestInvoke_MissingProvider(t *testing.T) { } func TestInvoke_MissingModel(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -196,8 +291,7 @@ func TestInvoke_MissingModel(t *testing.T) { } func TestInvoke_MissingTokens(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) cases := []struct { name string @@ -240,8 +334,7 @@ func TestInvoke_MissingTokens(t *testing.T) { } func TestInvoke_UnparseableTokens(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) cases := []struct { name string @@ -271,8 +364,7 @@ func TestInvoke_UnparseableTokens(t *testing.T) { } func TestInvoke_ZeroTokens(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -291,8 +383,7 @@ func TestInvoke_ZeroTokens(t *testing.T) { } func TestInvoke_UnknownModel(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -309,8 +400,7 @@ func TestInvoke_UnknownModel(t *testing.T) { } func TestInvoke_NilInput(t *testing.T) { - configureBuiltin(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfig(t)) out, err := mw.Invoke(context.Background(), nil) require.NoError(t, err) @@ -319,36 +409,12 @@ func TestInvoke_NilInput(t *testing.T) { assert.Empty(t, out.Metadata, "no metadata must be emitted on nil input") } -const fixturePricingWithCache = `openai: - gpt-4o: - input_per_1k: 0.0025 - output_per_1k: 0.01 - cached_input_per_1k: 0.00125 -anthropic: - claude-sonnet-4-5: - input_per_1k: 0.003 - output_per_1k: 0.015 - cache_read_per_1k: 0.0003 - cache_creation_per_1k: 0.00375 -` - -// configureBuiltinWithCacheRates points the package-level -// FactoryContext at a tmp directory containing pricing entries that -// include the cache rate fields. -func configureBuiltinWithCacheRates(t *testing.T) { - t.Helper() - dir := t.TempDir() - require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricingWithCache), 0o600), "write cache-aware pricing fixture") - builtin.Configure(context.Background(), dir, nil, nil, nil) -} - // TestInvoke_OpenAICachedSubsetDiscount proves the OpenAI shape end // to end through the middleware: cached_input_tokens is treated as a // SUBSET of input_tokens and discounted at the configured rate, not // added on top. func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -390,8 +456,7 @@ func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) { // shape: cache_read and cache_creation are additive to input_tokens // and each carries its own rate. func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -429,8 +494,37 @@ func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) { "output bucket bills 200 tokens at 0.015/1k") } +// TestInvoke_PerRecordEntryUsesSurfaceFormula: a per-record entry for an +// anthropic-surface request must bill its cache buckets additively — the +// formula follows llm.provider, not which table the entry came from. +func TestInvoke_PerRecordEntryUsesSurfaceFormula(t *testing.T) { + raw, err := json.Marshal(Config{Pricing: &PricingConfig{ + Providers: map[string]map[string]pricing.EntryJSON{ + "prov-ant": {"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375}}, + }, + }}) + require.NoError(t, err) + mw := buildMiddleware(t, raw) + + out, err := mw.Invoke(context.Background(), &middleware.Input{ + Metadata: []middleware.KV{ + {Key: middleware.KeyLLMProvider, Value: "anthropic"}, + {Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"}, + {Key: middleware.KeyLLMResolvedProviderID, Value: "prov-ant"}, + {Key: middleware.KeyLLMInputTokens, Value: "256"}, + {Key: middleware.KeyLLMOutputTokens, Value: "200"}, + {Key: middleware.KeyLLMCachedInputTokens, Value: "768"}, + {Key: middleware.KeyLLMCacheCreationTokens, Value: "512"}, + }, + }) + require.NoError(t, err) + value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) + require.True(t, ok) + assert.Equal(t, "0.005918400", value, "identical math to the defaults-table entry with the same rates") +} + // assertBucket asserts one per-bucket cost key carries the expected -// 6-decimal value. +// 9-decimal value. func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) { t.Helper() got, ok := metaValue(t, md, key) @@ -439,13 +533,10 @@ func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) { } // TestInvoke_CachedTokensAbsentFallsBackToBaseFormula covers the -// "operator hasn't opted in" path: with no cached metadata keys -// emitted, the meter must produce exactly the same cost as before -// the feature landed. Critical so operators with the new binary but -// no YAML changes see no behavioural drift on OpenAI requests. +// no-cache-metadata path: with no cached keys emitted, the meter must +// produce exactly the input+output cost. func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -460,7 +551,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) { value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal) require.True(t, ok) // 1000 input * 0.0025 + 500 output * 0.01 = 0.0025 + 0.005 = 0.0075 - assert.Equal(t, "0.007500000", value, "no cached metadata = same cost as before the feature landed") + assert.Equal(t, "0.007500000", value, "no cached metadata = plain input+output cost") } // TestInvoke_UnparseableCachedTokensSkippedSilently proves the @@ -469,8 +560,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) { // regular formula. Cache buckets are a refinement, never a reason to // abort cost computation. func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) { - configureBuiltinWithCacheRates(t) - mw := buildMiddleware(t, nil) + mw := buildMiddleware(t, fixtureConfigWithCache(t)) out, err := mw.Invoke(context.Background(), &middleware.Input{ Metadata: []middleware.KV{ @@ -487,22 +577,10 @@ func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) { assert.Equal(t, "0.007500000", value, "same as the no-cached-metadata path") } -// TestMiddleware_CloseCancelsReloader proves Close stops the per-instance -// pricing-reload goroutine: a chain rebuild retires the old instance and -// calls Close, which must invoke the cancel func startReloader handed it so -// the mtime-poll loop doesn't outlive the chain. -func TestMiddleware_CloseCancelsReloader(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - m := newMiddleware(nil, cancel) - - require.NoError(t, m.Close(), "Close must not error") - require.Error(t, ctx.Err(), "Close must cancel the reloader context so the poll goroutine exits") -} - -// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) for an -// instance with no reloader and for a nil receiver. +// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) even +// for a nil receiver. func TestMiddleware_CloseNilSafe(t *testing.T) { - require.NoError(t, newMiddleware(nil, nil).Close(), "no-reloader Close must be a no-op") + require.NoError(t, newMiddleware(nil, nil).Close(), "Close must be a no-op") var m *Middleware require.NoError(t, m.Close(), "nil-receiver Close must be safe") } diff --git a/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go b/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go index 827b81d07..d8cd81437 100644 --- a/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go +++ b/proxy/internal/middleware/builtin/llm_request_parser/bedrock_test.go @@ -6,23 +6,6 @@ import ( "github.com/stretchr/testify/require" ) -func TestNormalizeBedrockModel(t *testing.T) { - cases := map[string]string{ - "eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", - "us.anthropic.claude-opus-4-8-20250101-v1:0": "anthropic.claude-opus-4-8", - "apac.anthropic.claude-haiku-4-5-v1:0": "anthropic.claude-haiku-4-5", - "anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", - "meta.llama3-3-70b-instruct-v1:0": "meta.llama3-3-70b-instruct", - "amazon.nova-pro-v1:0": "amazon.nova-pro", - "amazon.nova-2-lite-v1:0": "amazon.nova-2-lite", - // Inference-profile ARN — model id lives in the last path segment. - "arn:aws:bedrock:eu-central-1:123456789012:inference-profile/eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", - } - for in, want := range cases { - require.Equal(t, want, normalizeBedrockModel(in), "normalize %q", in) - } -} - func TestParseBedrockPath(t *testing.T) { tests := []struct { path string diff --git a/proxy/internal/middleware/builtin/llm_request_parser/middleware.go b/proxy/internal/middleware/builtin/llm_request_parser/middleware.go index 64ca04e6a..b4d1e16d4 100644 --- a/proxy/internal/middleware/builtin/llm_request_parser/middleware.go +++ b/proxy/internal/middleware/builtin/llm_request_parser/middleware.go @@ -8,7 +8,6 @@ package llm_request_parser import ( "context" "net/url" - "regexp" "strconv" "strings" "unicode/utf8" @@ -253,9 +252,7 @@ func parseVertexPath(reqPath string) (vertexRequest, bool) { if c := strings.LastIndex(rest, ":"); c >= 0 { model, action = rest[:c], rest[c+1:] } - if at := strings.Index(model, "@"); at >= 0 { - model = model[:at] - } + model = llm.NormalizeVertexModel(model) if model == "" { return vertexRequest{}, false } @@ -343,14 +340,6 @@ func trimBedrockNamespace(reqPath string) string { return reqPath } -// bedrockRegionPrefixes are the cross-region inference-profile prefixes that -// front a Bedrock model id (e.g. "eu.anthropic.claude-..."). -var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."} - -// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]" -// version/throughput suffix of a Bedrock model id. -var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`) - // parseBedrockPath extracts the model and streaming/converse flags from an AWS // Bedrock runtime model endpoint: // @@ -375,7 +364,7 @@ func parseBedrockPath(reqPath string) (bedrockRequest, bool) { if decoded, err := url.PathUnescape(rawModel); err == nil { rawModel = decoded } - model := normalizeBedrockModel(rawModel) + model := llm.NormalizeBedrockModel(rawModel) if model == "" { return bedrockRequest{}, false } @@ -389,30 +378,6 @@ func parseBedrockPath(reqPath string) (bedrockRequest, bool) { } } -// normalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile -// prefix, and the version/throughput suffix from a Bedrock model id so it -// matches the catalog/pricing key, e.g. -// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5" -// and "arn:aws:bedrock:eu-central-1:123:inference-profile/eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -// -> "anthropic.claude-sonnet-4-5". -func normalizeBedrockModel(modelID string) string { - m := modelID - // A full ARN (inference-profile / provisioned-throughput / foundation-model) - // carries the model id in its last path segment. - if strings.HasPrefix(m, "arn:") { - if i := strings.LastIndex(m, "/"); i >= 0 { - m = m[i+1:] - } - } - for _, p := range bedrockRegionPrefixes { - if strings.HasPrefix(m, p) { - m = m[len(p):] - break - } - } - return bedrockVersionSuffix.ReplaceAllString(m, "") -} - // invokeBedrock emits the model/provider/session/prompt for an AWS Bedrock // request. Bedrock is metered under the dedicated "bedrock" parser, which reads // both the InvokeModel and Converse response shapes. diff --git a/proxy/internal/proxy/agent_network_chain_realstack_test.go b/proxy/internal/proxy/agent_network_chain_realstack_test.go index bc611fc98..924d37ace 100644 --- a/proxy/internal/proxy/agent_network_chain_realstack_test.go +++ b/proxy/internal/proxy/agent_network_chain_realstack_test.go @@ -18,10 +18,10 @@ import ( "google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/test/bufconn" - rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" - mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/proxy/internal/middleware" "github.com/netbirdio/netbird/proxy/internal/middleware/bodytap" @@ -134,14 +134,18 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) { RedactPii: true, })) require.NoError(t, st.SaveAgentNetworkProvider(ctx, &agentNetworkTypes.Provider{ - ID: providerID, - AccountID: testAccountID, - ProviderID: "openai_api", - Name: "openai-fullchain-test", - UpstreamURL: upstream.URL, // router rewrites to this - APIKey: "sk-test", - Enabled: true, - Models: []agentNetworkTypes.ProviderModel{{ID: "gpt-5.4"}}, + ID: providerID, + AccountID: testAccountID, + ProviderID: "openai_api", + Name: "openai-fullchain-test", + UpstreamURL: upstream.URL, // router rewrites to this + APIKey: "sk-test", + Enabled: true, + // Operator-pinned prices deliberately differ from the catalog's + // gpt-5.4 rates (0.0025/0.015) so the cost assertion below proves + // the per-provider-record price — not the default table — billed + // this request: stored price → synth → wire → cost_meter. + Models: []agentNetworkTypes.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}}, SessionPrivateKey: "priv", SessionPublicKey: "pub", })) @@ -176,7 +180,7 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) { // ---- 5. Wire the middleware framework — same registry the proxy uses // in production, configured with our bufconn-backed management client. - mwbuiltin.Configure(ctx, t.TempDir(), nil, testLogger, mgmtClient) + mwbuiltin.Configure(ctx, nil, testLogger, mgmtClient) registry := mwbuiltin.DefaultRegistry() mwMetrics, err := middleware.NewMetrics(nil) require.NoError(t, err) @@ -283,13 +287,23 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) { if r.DimensionKind == agentNetworkTypes.DimensionGroup && r.DimensionID == adminGroupID && r.WindowSeconds == 60 && - r.TokensInput+r.TokensOutput > 0 { + r.TokensInput+r.TokensOutput > 0 && + r.CostUSD > 0 { return true } } return false }, 5*time.Second, 50*time.Millisecond, - "Admins group consumption row must increment via the response leg — if this fails the proxy's respInput dropped UserGroups again or the parser/recorder wiring is broken") + "Admins group consumption row must increment via the response leg WITH a non-zero cost — a zero cost means the operator's stored price never reached cost_meter (synth → wire → per-record lookup broken)") + + // 8a-cost. Exact cost from the OPERATOR's stored price, not the catalog + // default: 12 prompt tokens × 0.004/1k + 40 completion tokens × 0.02/1k + // = 0.000048 + 0.0008 = 0.000848. With catalog rates it would be 0.00063 + // — this assertion distinguishes the two, closing the loop on the whole + // dynamic-pricing feature (dashboard save → synth → gRPC wire → + // llm.resolved_provider_id lookup → billing). + assert.Equal(t, "0.000848000", cd.GetMetadata()["cost.usd_total"], + "cost must be computed from the provider record's operator-pinned price") // 8b. Both the captured prompt and the captured completion are // redacted — proves the synth threads redact_pii=true into BOTH parser diff --git a/proxy/server.go b/proxy/server.go index 4f448e4b8..bd70b7e70 100644 --- a/proxy/server.go +++ b/proxy/server.go @@ -246,10 +246,6 @@ type Server struct { // in processMappings before the receive loop reconnects to resync. // Zero uses defaultMappingBatchWatchdog. MappingBatchWatchdog time.Duration - // MiddlewareDataDir is the base directory the middleware system uses to - // resolve file-backed configuration (e.g. the cost_meter pricing table). - // Empty means any middleware that requires a file fails at configure time. - MiddlewareDataDir string // MiddlewareCaptureBudgetBytes overrides the proxy-wide in-flight capture // budget passed to middleware.NewManager. Zero or negative values fall // back to defaultMiddlewareCaptureBudgetBytes (256 MiB). @@ -2093,7 +2089,7 @@ func (s *Server) initMiddlewareManager(ctx context.Context) error { return fmt.Errorf("middleware manager requires metrics bundle") } otelMeter := s.meter.Meter() - mwbuiltin.Configure(ctx, s.MiddlewareDataDir, otelMeter, s.Logger, s.mgmtClient) + mwbuiltin.Configure(ctx, otelMeter, s.Logger, s.mgmtClient) mwMetrics, err := middleware.NewMetrics(otelMeter) if err != nil { diff --git a/shared/llm/model.go b/shared/llm/model.go new file mode 100644 index 000000000..08e42e5a4 --- /dev/null +++ b/shared/llm/model.go @@ -0,0 +1,58 @@ +// Package llm holds LLM model-identifier helpers shared by the proxy and +// the management server. The proxy normalizes model ids parsed off inbound +// requests; management normalizes the operator's registered model ids at +// synthesis time so both sides of the pricing / routing contract compare +// equal. +package llm + +import ( + "regexp" + "strings" +) + +// bedrockRegionPrefixes are the cross-region inference-profile prefixes that +// front a Bedrock model id (e.g. "eu.anthropic.claude-..."). +var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."} + +// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]" +// version/throughput suffix of a Bedrock model id. +var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`) + +// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile +// prefix, and the version/throughput suffix from a Bedrock model id so it +// matches the catalog/pricing key, e.g. +// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5" +// and the inference-profile ARN's last segment likewise. It is the single +// source of truth shared by the proxy's request parser (which normalizes the +// request model from the URL path), the proxy's router (which normalizes the +// operator's registered Bedrock model ids so both sides compare equal), and +// the management synthesizer (which keys per-provider pricing entries by the +// normalized id the parser will emit at billing time). +func NormalizeBedrockModel(modelID string) string { + m := modelID + // A full ARN (inference-profile / provisioned-throughput / foundation-model) + // carries the model id in its last path segment. + if strings.HasPrefix(m, "arn:") { + if i := strings.LastIndex(m, "/"); i >= 0 { + m = m[i+1:] + } + } + for _, p := range bedrockRegionPrefixes { + if strings.HasPrefix(m, p) { + m = m[len(p):] + break + } + } + return bedrockVersionSuffix.ReplaceAllString(m, "") +} + +// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id +// (e.g. "claude-sonnet-4-5@20250929" -> "claude-sonnet-4-5") so it matches +// the catalog/pricing key. Vertex publisher models are priced under their +// vendor surface with the bare, unversioned id. +func NormalizeVertexModel(modelID string) string { + if at := strings.Index(modelID, "@"); at >= 0 { + return modelID[:at] + } + return modelID +} diff --git a/proxy/internal/llm/bedrock_model_test.go b/shared/llm/model_test.go similarity index 66% rename from proxy/internal/llm/bedrock_model_test.go rename to shared/llm/model_test.go index 3bd9662b7..42f2e9ca5 100644 --- a/proxy/internal/llm/bedrock_model_test.go +++ b/shared/llm/model_test.go @@ -11,6 +11,8 @@ func TestNormalizeBedrockModel(t *testing.T) { "eu.anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", "us.anthropic.claude-haiku-4-5": "anthropic.claude-haiku-4-5", "us.anthropic.claude-opus-4-8-20250101-v1:0": "anthropic.claude-opus-4-8", + "apac.anthropic.claude-haiku-4-5-v1:0": "anthropic.claude-haiku-4-5", + "amazon.nova-2-lite-v1:0": "amazon.nova-2-lite", "anthropic.claude-sonnet-4-5-20250929-v1:0": "anthropic.claude-sonnet-4-5", "meta.llama3-3-70b-instruct-v1:0": "meta.llama3-3-70b-instruct", "amazon.nova-pro-v1:0": "amazon.nova-pro", @@ -21,3 +23,14 @@ func TestNormalizeBedrockModel(t *testing.T) { require.Equal(t, want, NormalizeBedrockModel(in), "normalize %q", in) } } + +func TestNormalizeVertexModel(t *testing.T) { + cases := map[string]string{ + "claude-sonnet-4-5@20250929": "claude-sonnet-4-5", + "claude-haiku-4-5": "claude-haiku-4-5", + "gpt-4o@2024-08-06": "gpt-4o", + } + for in, want := range cases { + require.Equal(t, want, NormalizeVertexModel(in), "normalize %q", in) + } +} diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index e3d11227a..8ad3d932c 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -5271,6 +5271,21 @@ components: format: double description: Cost per 1k output tokens, in USD. example: 0.0006 + cached_input_per_1k: + type: number + format: double + description: OpenAI-shape cache rate — cost per 1k cached prompt tokens (a subset of input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means no discount (cached tokens bill at input_per_1k). + example: 0.000075 + cache_read_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — cost per 1k cache-read tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache reads bill at input_per_1k. + example: 0.0003 + cache_creation_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — cost per 1k cache-creation tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache writes bill at input_per_1k. + example: 0.00375 required: - id - input_per_1k @@ -5296,6 +5311,21 @@ components: format: double description: Output token price per 1k tokens, in USD. example: 0.015 + cached_input_per_1k: + type: number + format: double + description: OpenAI-shape cache rate — default cost per 1k cached prompt tokens (a subset of input tokens), in USD. Absent when the model has no cached-input discount. + example: 0.000075 + cache_read_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — default cost per 1k cache-read tokens (additive to input tokens), in USD. Absent when the model has no cache-read rate. + example: 0.0003 + cache_creation_per_1k: + type: number + format: double + description: Anthropic-shape cache rate — default cost per 1k cache-creation tokens (additive to input tokens), in USD. Absent when the model has no cache-creation rate. + example: 0.00375 context_window: type: integer description: Maximum context window in tokens. @@ -5354,6 +5384,13 @@ components: $ref: '#/components/schemas/AgentNetworkCatalogExtraHeader' identity_injection: $ref: '#/components/schemas/AgentNetworkCatalogIdentityInjection' + pricing_surfaces: + type: array + description: | + Cost-meter pricing surfaces this provider's traffic is metered under ("openai", "anthropic", "bedrock"). Tells the dashboard which cache-rate fields apply to this provider's models: "openai" → cached_input_per_1k (cached prompt tokens are a subset of input); "anthropic"/"bedrock" → cache_read_per_1k + cache_creation_per_1k (additive buckets). Absent/empty for gateway and custom entries, whose upstream shape NetBird cannot know ahead of time — surface all cache fields for those. + items: + type: string + example: ["openai"] models: type: array description: Catalog models available for this provider. diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index a4de48a09..87dd9ccfc 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -2017,6 +2017,15 @@ type AgentNetworkCatalogJSONMetadataInjection struct { // AgentNetworkCatalogModel defines model for AgentNetworkCatalogModel. type AgentNetworkCatalogModel struct { + // CacheCreationPer1k Anthropic-shape cache rate — default cost per 1k cache-creation tokens (additive to input tokens), in USD. Absent when the model has no cache-creation rate. + CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"` + + // CacheReadPer1k Anthropic-shape cache rate — default cost per 1k cache-read tokens (additive to input tokens), in USD. Absent when the model has no cache-read rate. + CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"` + + // CachedInputPer1k OpenAI-shape cache rate — default cost per 1k cached prompt tokens (a subset of input tokens), in USD. Absent when the model has no cached-input discount. + CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"` + // ContextWindow Maximum context window in tokens. ContextWindow int `json:"context_window"` @@ -2070,6 +2079,9 @@ type AgentNetworkCatalogProvider struct { // Name Display name for the provider. Name string `json:"name"` + + // PricingSurfaces Cost-meter pricing surfaces this provider's traffic is metered under ("openai", "anthropic", "bedrock"). Tells the dashboard which cache-rate fields apply to this provider's models: "openai" → cached_input_per_1k (cached prompt tokens are a subset of input); "anthropic"/"bedrock" → cache_read_per_1k + cache_creation_per_1k (additive buckets). Absent/empty for gateway and custom entries, whose upstream shape NetBird cannot know ahead of time — surface all cache fields for those. + PricingSurfaces *[]string `json:"pricing_surfaces,omitempty"` } // AgentNetworkCatalogProviderKind Presentation grouping for the provider Select on the dashboard. @@ -2293,6 +2305,15 @@ type AgentNetworkProvider struct { // AgentNetworkProviderModel A model exposed by the provider, with the operator's per-1k input/output prices in USD. type AgentNetworkProviderModel struct { + // CacheCreationPer1k Anthropic-shape cache rate — cost per 1k cache-creation tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache writes bill at input_per_1k. + CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"` + + // CacheReadPer1k Anthropic-shape cache rate — cost per 1k cache-read tokens (additive to input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means cache reads bill at input_per_1k. + CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"` + + // CachedInputPer1k OpenAI-shape cache rate — cost per 1k cached prompt tokens (a subset of input tokens), in USD. Omitted means inherit NetBird's default rate for this model when one exists; 0 means no discount (cached tokens bill at input_per_1k). + CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"` + // Id Model identifier (e.g. "gpt-4o-mini"). Id string `json:"id"`