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 01/34] [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 02/34] [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 03/34] [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 04/34] [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 05/34] [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 06/34] [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 07/34] [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 08/34] [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 09/34] [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"` From e56eb14c521229a6f0416c5db440e115596dcee5 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Sat, 1 Aug 2026 08:17:03 +0900 Subject: [PATCH 10/34] [misc] add AGENTS.md file (#7014) ## Describe your changes [#7009](https://github.com/netbirdio/netbird/pull/7009) wrote down what we expect from a contribution: an agreed ticket first, a change the author has run, one purpose per PR, small enough to review, a title tag CI already enforces. That works for contributors who read CONTRIBUTING.md. A growing share of what reaches the queue is drafted by a coding agent, and those tools never read it. The result is predictable and repetitive: a PR with no ticket behind it, an approach we would have talked the author out of in five minutes, a diff too large to review carefully against a daemon that runs as root, a description longer than the code it describes, and an author who cannot answer questions about their own change. That is not a tooling problem. It is what happens when a capable tool is pointed at a repository whose expectations nobody told it. AGENTS.md states those expectations in the format agents pick up automatically. The first entry in its stop-and-ask list is asking the contributor for the discussion or issue before drafting anything, which is where most of these PRs go wrong. It also carries the repository map, the Go conventions we apply in review, the local verification commands, the PR template and title-tag rules, and instructions to read the review bots rather than resolve their threads. CONTRIBUTING.md gains a short section saying plainly that we have no policy for or against agents, that this guide exists because of what we keep seeing rather than as a restriction on tools, and that whatever produced a diff its author owns it. It also records that we weigh whether a contribution is worth maintaining, and that what gets merged has to match our security and design expectations. CLAUDE.md is a one-line pointer rather than a symlink, deliberately. A symlink is tidier, but a Windows checkout without symlink support materialises it as a plain file containing the target path, and an agent then reads one word as its entire instruction set with no error to show for it. Given how much Windows work this repository sees, a pointer file that survives every checkout is worth the extra file. Nothing here is enforced by CI, and no workflow changes. It is guidance, aimed at PRs arriving in a reviewable state more often. --- AGENTS.md | 514 ++++++++++++++++++++++++++++++++++++++++++++++++ CLAUDE.md | 1 + CONTRIBUTING.md | 30 +++ 3 files changed, 545 insertions(+) create mode 100644 AGENTS.md create mode 100644 CLAUDE.md diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..4ac006795 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,514 @@ +# NetBird Agent Guidelines + +**NetBird** is an open-source connectivity platform: a WireGuard®-based overlay +network with a control plane. The **agent** (`client/`) runs on user machines as +a privileged daemon and manages the WireGuard interface, routing, firewall, and +DNS. **Management** (`management/`) is the control plane and REST/gRPC API, +**Signal** (`signal/`) brokers peer handshakes, **Relay** (`relay/`) carries +traffic when a direct tunnel is impossible, and **Proxy** (`proxy/`) is the +identity-aware proxy behind Agent Network. + +This file applies to the whole repository, and is the single source of truth for +agent guidance here. `CLAUDE.md` is a one-line pointer to it — keep the guidance +in this file, not duplicated there. + +## Contents + +- [NetBird Agent Guidelines](#netbird-agent-guidelines) + - [Contents](#contents) + - [STOP and ask the user before](#stop-and-ask-the-user-before) + - [Quick reference](#quick-reference) + - [Structure](#structure) + - [Where to look](#where-to-look) + - [Repo-wide principles](#repo-wide-principles) + - [Error handling](#error-handling) + - [Comments](#comments) + - [Testing](#testing) + - [Pitfalls](#pitfalls) + - [Commits, PRs, releases](#commits-prs-releases) + - [After you push: CI and review bots](#after-you-push-ci-and-review-bots) + - [Discussion and support](#discussion-and-support) + +## STOP and ask the user before + +- **Opening a pull request for anything beyond a trivial fix, without an agreed + ticket.** Ask the user directly: *"Is there a discussion or issue for this + change?"* NetBird is discussion-first — community reports start in + [Discussions](https://github.com/netbirdio/netbird/discussions), DevRel + validates them, and only validated discussions become issues. A PR that + changes behavior with no linked issue may be closed on arrival. If there is no + ticket, offer to draft the discussion post **instead of** the PR, and wait for + the user's call. Only typos, broken links, documentation corrections, and + one-line fixes that already have an issue can skip this. +- **Designing in any high-risk area** (see + [CONTRIBUTING.md](CONTRIBUTING.md#high-risk-areas)): public API and OpenAPI + schema, gRPC protos, behavior existing deployments would notice after an + upgrade, peer connectivity (ICE, NAT traversal, relay selection, WireGuard® or + Rosenpass key handling), client system integration (routing, firewall, DNS, + interface), authentication and authorization, CLI or service flags, config + file format, daemon IPC, store schema and migrations, or a new feature. The + design gets agreed in the ticket before code is written. +- **Writing a store migration or changing a persisted model.** Migrations are + one-way in the field and both the GORM and pgx paths may need the change. +- **Hand-editing generated code.** `*.pb.go`, `*.gen.go`, and mocks are outputs. + Edit the source (`.proto`, `openapi.yml`) and rerun the matching + `generate.sh`. +- **Adding, removing, or bumping a dependency**, and never vendor a fork. +- **Weakening a security control** — authentication, authorization, certificate + verification, privilege dropping, or peer identity checks — even when it is + the fastest way to make a test pass. +- **Force-pushing to `main`**, force-pushing any branch that is already under + review, amending pushed commits, or bypassing hooks with `--no-verify`. + +## Quick reference + +```bash +# Build +go build ./... +cd client && CGO_ENABLED=0 go build . # agent +cd management && go build . # management service +cd signal && go build . # signal service + +# Verify (run before every push) +go fmt ./... +make lint # golangci-lint on files changed vs origin/main (also the pre-push hook) +make lint-all # full-repository lint, matches CI +make test-unit # host-safe unit tests, -tags devcert, no sudo +make test-privileged # privileged-tagged suite in a Docker container with NET_ADMIN +make setup-hooks # wire make lint into .githooks/pre-push + +# Narrow runs +go test ./client/internal/dns/... +go test -race -run TestPeerConn ./client/internal/peer/... +PRIV_RUN=TestNftablesManager PRIV_PKGS=./client/firewall/nftables/... make test-privileged + +# Code generation (never hand-edit the output) +./shared/management/http/api/generate.sh # REST types from openapi.yml +./shared/management/proto/generate.sh +./shared/signal/proto/generate.sh +./client/proto/generate.sh +./flow/proto/generate.sh + +# Run locally (lab only, never on a machine you rely on) +sudo ./client/netbird up --log-level debug --log-file console +sudo ./client/netbird down # teardown: restores routing, firewall, DNS +./signal/signal run --log-level debug --log-file console +./management/management management --log-level debug --log-file console --config ./management.json +``` + +`netbird up` needs root and rewrites the host's routing table, firewall rules, +DNS configuration, and WireGuard® interface. Run it only in a disposable test +environment (a VM, container, or throwaway host) that you can rebuild, never on +a workstation or server whose connectivity matters. Run `sudo netbird down` +before you stop working, before rebuilding the binary, and on every failure +path, so the host's networking state is restored instead of left half-applied. +See [Pitfalls](#pitfalls) for why cleanup on every exit path matters. + +## Structure + +```text +netbird/ +├── client/ NetBird agent +│ ├── cmd/ agent CLI +│ ├── internal/ agent business logic (engine, peer, dns, routemanager, ...) +│ ├── server/ daemon for background execution +│ ├── proto/ daemon gRPC protos +│ ├── iface/ WireGuard® interface management +│ ├── firewall/ nftables, iptables, pf, WFP, userspace backends +│ ├── ssh/ built-in SSH server and client +│ ├── ui/ desktop UI (Wails v3 + React) +│ ├── android/, ios/ mobile bindings +│ ├── wasm/ WebAssembly build +│ └── mdm/, system/ MDM policy, host information +├── management/ control plane +│ └── server/ account, peer, groups, networks, posture, permissions, +│ settings, store, http (REST), idp, integrations, migration +├── signal/ handshake broker (peer/, server/) +├── relay/ relay service (protocol/, server/, healthcheck/) +├── proxy/ identity-aware proxy (llm/, acme/, accesslog/, middleware/, tcp/, udp/) +├── agent-network/ Agent Network overview +├── shared/ imported by both agent and services +│ ├── management/ proto/, client/, http/api (OpenAPI + generated types) +│ ├── signal/ proto/, client/ +│ └── relay/, auth/, sshauth/, metrics/ +├── e2e/ end-to-end suites and harness +├── encryption/, dns/, route/, stun/, sharedsock/, util/, flow/ +├── infrastructure_files/ docker compose and getting-started templates +└── release_files/ files packaged into releases +``` + +## Where to look + +| Task | Location | +| --------------------------- | ------------------------------------------------------------ | +| REST API / OpenAPI | `shared/management/http/api/` + `management/server/http/` | +| Management gRPC protocol | `shared/management/proto/` | +| Signal protocol | `shared/signal/proto/` | +| Daemon IPC protocol | `client/proto/` | +| Peer connection and NAT | `client/internal/peer/` | +| Network map handling | `client/internal/engine.go`, `shared/management/networkmap/` | +| Routing | `client/internal/routemanager/`, `route/` | +| Firewall backends | `client/firewall/` | +| DNS | `client/internal/dns/`, `dns/` | +| WireGuard® interface | `client/iface/` | +| Persistence and migrations | `management/server/store/`, `management/server/migration/` | +| IdP integrations | `management/server/idp/` | +| Permissions model | `management/server/permissions/` | +| LLM routing / Agent Network | `proxy/internal/llm/`, `agent-network/` | +| End-to-end tests | `e2e/` | + +## Repo-wide principles + +1. **Run `go fmt` on every modified Go file.** Formatting is not optional. +2. **Zero unaddressed diagnostics.** Fix IDE and linter warnings on code you + touch, and delete imports, helpers, and parameters your refactor orphaned. + Exception: unused parameters in shared code may be consumed by builds outside + this repository — do not remove them, ask instead. +3. **Function comments are mandatory for exported functions**, written as full + sentences with a period, starting with the identifier name. +4. **Prefer private functions and constants.** Export only what a caller outside + the package genuinely needs. +5. **Early returns and guard clauses.** Handle errors and edge cases first + instead of nesting `if`/`else` chains. +6. **Split complex functions.** If a function trips a complexity warning, break + it into named helpers rather than silencing the warning. +7. **Avoid LLM-slop tells:** em dashes, hedging narration, restating the diff in + prose, trailing summaries. Defaults, not absolute bans. Applies to code, + comments, commit messages, and PR descriptions alike. +8. **Concurrency: do a two-pass race analysis after every change** that adds + shared state. Guard maps and slices with a mutex, keep critical sections + short, and run `go test -race` on the touched packages. +9. **Cross-platform builds must keep working.** The agent targets Linux, macOS, + Windows, FreeBSD, Android, and iOS. When you add a platform-specific file, + add the counterpart or a build-tagged fallback for the others. +10. **Never hand-edit generated files.** Change the source and regenerate. +11. **Never log secrets** — private keys, setup keys, tokens, PAT values — and + keep peer IPs and hostnames out of logs above debug level. + +## Error handling + +Use single-assignment form when the error is only needed inside the `if`: + +```go +// Good +if err := someCall(); err != nil { + return fmt.Errorf("context: %w", err) +} + +// Bad - unnecessary split +err := someCall() +if err != nil { + return fmt.Errorf("context: %w", err) +} +``` + +Use multiple assignment when the value is needed after the block: + +```go +result, err := someCall() +if err != nil { + return fmt.Errorf("context: %w", err) +} +``` + +Add short, meaningful context, and **do not** start `fmt.Errorf` messages with +obvious words like "failed to" or "error": + +```go +// Good +return fmt.Errorf("parse remote address: %w", err) +return fmt.Errorf("listen on %s: %w", addr, err) + +// Bad +return fmt.Errorf("failed to parse remote address: %w", err) +return fmt.Errorf("error listening on %s: %w", addr, err) + +// "failed" is fine in log messages +log.Debugf("failed to parse remote address: %v", err) +``` + +Skip the wrapping when a function only extracts or delegates and the wrap would +add nothing: + +```go +func parseAddr(addr string) (string, int, error) { + host, portStr, err := net.SplitHostPort(addr) + if err != nil { + return "", 0, err + } + // ... +} +``` + +Log the errors you choose not to act on: + +- `log.Debugf()` for errors that do not affect program flow but help debugging. +- `log.Tracef()` for very verbose errors that would otherwise spam logs. +- **Never ignore** errors from writes, network sends, or critical cleanup. +- Close errors may be ignored for read-only operations; log them at debug for + writes. + +## Comments + +Comment the **why**, never the **what**. Default to no comment, and add one only +when a hidden constraint or workaround would surprise a future reader. Never +reference the current task, PR, or your own changes in a comment. + +```go +// Bad - trailing comments explaining the obvious +defer localConn.Close() // Close the connection +if err != nil { // Check if error occurred + +// Good +defer localConn.Close() + +// Good - explains a non-obvious constraint +// Use incremental checksum update per RFC 1624 for performance. +checksum = updateChecksum(checksum, oldPort, newPort) +``` + +### Length budget + +- **90 characters per line.** Wrap the comment, do not run past it. +- **250 characters per comment**, roughly three wrapped lines. Doc comments on + exported identifiers may exceed it when the API genuinely needs the + explanation; inline comments inside a function body may not. + +The budget is a smell detector, not a rule to game. Do not compress a needed +explanation into cryptic shorthand to fit — if a block of code needs more than +250 characters of prose, the code is doing too much. Fix the code: + +- **Extract a named function.** A well-named function replaces the comment: the + name says *what*, the body shows *how*, and the comment you no longer write + was the *what* anyway. Clean Code calls this "explain yourself in code". +- **Extract a named constant or predicate.** `if isExpiredSetupKey(key)` needs + no comment; `if key.ExpiresAt.Before(now) && !key.Revoked && key.UsageLimit > 0` + does. +- **Keep the surviving comment for the why** — the RFC, the kernel quirk, the + ordering constraint. That part is usually one or two lines. + +### Long switch and if/else chains + +A `switch` whose cases carry multi-line explanations is the usual place this +budget is breached, and the comment is a symptom. In order of preference: + +1. **Extract each case body into a named function.** The case becomes one line, + the name carries the meaning, and the switch reads as a table of contents. +2. **Replace the switch with a lookup table** — `map[Kind]handlerFunc` — when the + branches are uniform. Adding a case stops meaning editing a growing function. +3. **Replace conditional with polymorphism** when branches vary by type and the + same switch shape starts appearing in more than one place. Clean Code's rule + of thumb: tolerate a switch statement if it appears **once**, is buried in a + factory that returns an interface, and no other switch dispatches on the same + type. A second switch over the same enum is the signal to introduce the + interface. + +Do not restructure a switch purely to satisfy the budget when the cases are one +line each and self-evident — a flat, boring `switch` over an enum is fine and +needs no comments at all. + +Explanatory comments in tests are welcome — they document the scenario being set +up, and the 250-character budget does not apply to them. + +## Testing + +- **Unit tests** live beside the code as `_test.go`. `make test-unit` runs the + host-safe set with `-tags devcert` and no sudo. +- **Privileged tests** carry the `privileged` build tag and mutate host + networking. They run through `make test-privileged`, inside a Docker container + with `NET_ADMIN`. Never bypass that harness by running them directly on the + host. +- **End-to-end suites** live in `e2e/` with a shared harness. +- **Test real behavior, not API existence.** Assert on the observable end state + a consumer would see — bytes that arrived, the packet after translation, the + row after the write — not merely that a method exists or returns an error. +- **Avoid mocks for code we own.** Exercise the real store, manager, or + controller and assert what the caller actually receives. +- **`require` for setup and preconditions, `assert` for the conditions under + test.** Use `require` whenever a later line would panic or be meaningless + otherwise. +- **Message guidance:** optional for `NoError`/`Error`; always give context for + comparison, boolean, and collection assertions. + +```go +server, err := StartTestServer() +require.NoError(t, err, "Test server setup must succeed") +defer server.Close() + +result, err := client.DoOperation() +assert.NoError(t, err) +assert.Equal(t, expectedResult, result, "Result should match expected") +``` + +## Pitfalls + +- **The agent runs as root.** Anything touching routing, firewall, DNS, or the + interface can take a user's machine off the network. Prefer a reversible + change and make sure cleanup runs on every exit path. +- **Management has two account loaders** (GORM and pgx). Adding a relation to an + account often means updating both, or it silently comes back empty in + production. +- **`go test ./...` without `-tags devcert` skips tests** that need the + development certificate. Use `make test-unit`. +- **`make lint` only checks the diff against `origin/main`.** CI runs + `make lint-all`; run it too before pushing a large change. +- **Protos are consumed by released clients.** An old agent must keep working + against a new Management, so fields are added, never renumbered or removed. +- **Windows requires the wintun driver**, and the daemon serves a named pipe + (`npipe://netbird`) rather than loopback TCP. Loopback TCP carries no caller + identity, so privileged operations are refused over it. + +## Commits, PRs, releases + +- **PR titles must start with a bracketed tag.** Before you propose a title, + **read [`.github/workflows/pr-title-check.yml`](.github/workflows/pr-title-check.yml) + and take the allowed tags from the `allowedTags` array in that file.** It is + the only source of truth, it changes as components are added, and the check + runs on every title edit — a tag that is not in that array is a red build. Do + not rely on a list memorized from anywhere else, including this file. + + ```text + [client] Authorize daemon IPC callers by their local identity + [management,client] Add MDM policy support + ``` + + Multiple tags are comma-separated inside one pair of brackets. Match the tag + to the component you actually changed, not to the one you read the most. + +- **Use the repository's PR template.** Fill in + [`.github/pull_request_template.md`](.github/pull_request_template.md) rather + than replacing it with your own summary: describe the change, link the issue, + tick the checklist honestly (including "ran locally" and "single purpose"), + and complete the documentation section. Do not tick a box you have not + verified, and do not delete rows that do not apply. + +- **Keep the PR description short.** Under 1000 words on top of the template's + own text, and usually far less — a few paragraphs. Reviewers read the diff; + the description exists to explain what the diff cannot say for itself. This is + well below what an agent will produce by default, so cut before you post. + +- **Body: why before what.** Lead with the problem and the reason for this + approach, then the shape of the change. No bullet list of files changed, no + per-function walkthrough, no restating the diff in prose, no trailing summary + section, no self-congratulatory closing line. + +- **No `Co-Authored-By` or tool-attribution trailers in the PR description**, + and none in commits either. Contributors own their contributions. Whatever + tooling produced the diff, the person opening the PR is its author: they have + read every line, they can explain why it works, they can answer review + questions without going back to a model, and they are accountable for the + consequences of merging it. Do not add a trailer, footer, or description line + that spreads that ownership onto a tool. + +- **Commit subjects follow the same `[scope] Subject` convention.** Keep the + subject short, and use the body for why before what. No bullet lists of files + changed. + +- **Push review fixes as separate commits.** The PR is squashed on merge, so + there is no reason to rewrite history mid-review; many small commits make the + re-review readable. + +- **Do not force-push a branch that is under review.** A force-push detaches + existing review comments from the lines they were written against, destroys + the "changes since your last review" diff a reviewer relies on, and discards + the CI history that showed which commit broke what. Add commits instead — + including for fixups and reverts. Force-push only when there is no + alternative: a rebase to clear a genuine conflict, or removing a secret or a + large binary that was committed by mistake. When you must, ask the user first, + then say so in a PR comment so reviewers know their anchors moved. Never + force-push `main`, and never force-push a branch you do not own. + +- **One PR, one purpose.** Split refactors out of fixes and fixes out of + features. + +- **Keep the PR small.** Size is the single strongest predictor of how long a PR + waits. Aim for **under ~400 changed lines across under ~20 files**; past + roughly **1000 lines or 50 files** a community PR is likely to be sent back to + be split, or left unreviewed until it is. Large PRs from outside the core team + may be blocked outright when the size was never agreed in the ticket — + reviewing a sprawling change against a privileged networking daemon is a + security risk in itself, not just a time cost. + + Judge the size by hand-written code: exclude generated output, `go.sum`, + vendored files, and test fixtures from the estimate, but do not use their + presence to argue a 3000-line PR is small. + + When a change genuinely cannot be small — a protocol migration, a + cross-component rename — agree the split in the ticket **before** writing + code, and land it as a sequence of PRs that each build, test, and make sense + on their own. Propose that split to the user rather than opening one large PR + and hoping. + +- **User-facing changes need a docs PR** in + [netbirdio/docs](https://github.com/netbirdio/docs), linked from the PR + description. + +## After you push: CI and review bots + +Opening the PR is not the end of the task. Watch the run, read what the bots +say, and drive the PR to green before you report the work as done. + +```bash +gh pr checks --watch # all checks, live +gh run view --log-failed # only the failing steps +gh pr view --comments # bot and human review comments +``` + +**Never report a change as finished while checks are pending or red**, and never +describe a red PR as passing. If you ran out of turn before CI finished, say +which checks were still running. + +### The checks + +- **Go tests** — `golang-test-{linux,darwin,windows,freebsd}.yml`, sharded per + component. A failure in a component you did not touch is usually a real + interaction, not noise; read the log before assuming flake. +- **golangci-lint** — `golangci-lint.yml` runs the full repository, while + `make lint` only checks your diff. A clean local lint does not guarantee green + CI on a large change. +- **PR Title Check** — `pr-title-check.yml`, see above. +- **Codecov** — uploaded from the Linux test workflow with per-component flags + (`unit,client`, `unit,management`, `unit,relay`, `unit,proxy`, `unit,signal`, + `integration,management`). Coverage on new code should not go backwards. Add + tests for the paths you introduced; do not adjust thresholds or exclude files + to clear the report. +- **CodeRabbit** — configured in [`.coderabbit.yaml`](.coderabbit.yaml): `chill` + profile, auto-review on every non-draft PR, TypeScript/JavaScript/SVG paths + filtered out. Chat auto-reply is on, so `@coderabbitai` in a comment reaches + it. +- **SonarCloud** — project `netbirdio_netbird`, quality gate on new code (bugs, + vulnerabilities, code smells, duplication, coverage). +- **Snyk** — dependency and code scanning. + +Sonar and Snyk report as GitHub App checks rather than workflows in this +repository, so their detail lives on the PR check, not in the Actions logs. + +### Handling bot findings + +- **Read every comment and act on it.** Either fix it, or reply with the reason + it does not apply. Do not bulk-resolve threads to clear the count, and do not + silently ignore a finding because the check is advisory. +- **Bots are frequently wrong here.** NetBird has privileged, platform-specific, + and concurrency-heavy code that static analysis reads poorly. A confident + CodeRabbit or Sonar comment can still be nonsense. Verify the claim against + the code before you change anything — never edit correct code just to silence + a bot. +- **Security findings get the opposite default.** For a Snyk or Sonar + vulnerability, or a CodeRabbit comment about authentication, authorization, + certificate verification, or key handling, assume it is real until you have + disproved it. Surface it to the user rather than dismissing it yourself. +- **A new vulnerable dependency is a stop.** Bumping or replacing dependencies + needs the user's decision, as above. +- **Never change a workflow, threshold, lint exclusion, or bot config to make a + check pass.** If a check is genuinely wrong, say so and let the user decide. +- **Do not paper over flakes with blind re-runs.** Identify the failure first. If + it is a known flake, name it; if you cannot tell, report it as unresolved + rather than re-running until it goes green. + +## Discussion and support + +- Discussions: +- Slack: +- Docs: +- Security: — never in public +- Contribution process: [CONTRIBUTING.md](CONTRIBUTING.md) diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 000000000..764f406be --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1 @@ +See [AGENTS.md](AGENTS.md) for the agent guidelines in this repository. diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 3b8017788..db5097a48 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -66,11 +66,41 @@ Typical bug fixes, internal refactors, documentation updates, and tests do not need a design discussion, but should still be tied to an issue so the work is visible and nobody duplicates it. +### Using AI coding agents + +We have no policy for or against using an AI agent to write NetBird code. That +choice is yours, and we are not going to interrogate anyone about their tools. + +What we do have is a lot of incoming contributions that were plainly drafted with +one, and enough experience reviewing them to see the same avoidable problems +again and again: no ticket behind the change, a diff far too large to review, a +description longer than the code it describes, an approach that was never going +to be accepted, and an author who cannot answer questions about their own PR. +None of that is caused by the tooling — it is what happens when a tool is pointed +at a repository whose expectations it has never been told. + +So rather than a rule, there is a guide. [AGENTS.md](AGENTS.md) restates the +expectations from this document in the form agents read automatically +(`CLAUDE.md` points to it), so pointing your tool at the repository is usually +enough. Among other things it tells the agent to ask you for the +discussion or issue before drafting a PR, to keep the change small and +single-purpose, to run the tests locally, to use this repository's PR template +and title tags, and to write a description a reviewer can get through. + +The guardrails are the point, and they are the same ones we apply to everyone: an +agreed ticket, a change you have actually run, a diff small enough to review with +care, and an author who can explain it. Whatever wrote the diff, you are its +author — you own every line you submit and the consequences of opening a PR with it. + +We may assess whether a contribution is maintainable and whether its merged code +aligns with our security standards and design expectations. + ## Contents - [Contributing to NetBird](#contributing-to-netbird) - [Ticket first, PR second](#ticket-first-pr-second) - [High-risk areas](#high-risk-areas) + - [Using AI coding agents](#using-ai-coding-agents) - [Contents](#contents) - [Code of conduct](#code-of-conduct) - [Directory structure](#directory-structure) From 77f7e9fc912c338cc8988c90ebb84bdb9fa968f6 Mon Sep 17 00:00:00 2001 From: Ben Date: Sun, 2 Aug 2026 09:23:37 +0200 Subject: [PATCH 11/34] [client] Handle interface lookup errors in iOS DNS index helper (#6999) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes `getInterfaceIndex` in the iOS upstream DNS resolver dereferenced the result of `net.InterfaceByName` before checking the error, so a missing interface (e.g. during teardown or renaming) caused a nil-pointer panic instead of a DNS client error. The helper now returns a wrapped error before touching the interface; the only caller, `GetClientPrivate`, already propagates the error. The helper moved to an un-build-tagged file so it can be unit-tested on host platforms while remaining available to the iOS build. Added a test covering the missing-interface path. Verified with the new host test, an iOS arm64 CGO compile, and `git diff --check`. ## Issue ticket number and link N/A ## Stack Standalone PR based on `main`. ### 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) - [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 (explain why) Internal crash fix in the iOS DNS path; no user-facing behavior or configuration changes. ### Docs PR URL (required if "docs added" is checked) N/A ## Summary by CodeRabbit * **Bug Fixes** * Improved handling of network interface lookup failures with clearer error messages that identify the affected interface. * Added validation for network interface lookups, including reliable error handling when an interface cannot be found. * **Tests** * Added coverage for both successful interface resolution and missing-interface scenarios. --- client/internal/dns/interface_index.go | 15 +++++++++ client/internal/dns/interface_index_test.go | 35 +++++++++++++++++++++ client/internal/dns/upstream_ios.go | 5 --- 3 files changed, 50 insertions(+), 5 deletions(-) create mode 100644 client/internal/dns/interface_index.go create mode 100644 client/internal/dns/interface_index_test.go diff --git a/client/internal/dns/interface_index.go b/client/internal/dns/interface_index.go new file mode 100644 index 000000000..9e7dca080 --- /dev/null +++ b/client/internal/dns/interface_index.go @@ -0,0 +1,15 @@ +package dns + +import ( + "fmt" + "net" +) + +func getInterfaceIndex(interfaceName string) (int, error) { + iface, err := net.InterfaceByName(interfaceName) + if err != nil { + return 0, fmt.Errorf("lookup interface %q: %w", interfaceName, err) + } + + return iface.Index, nil +} diff --git a/client/internal/dns/interface_index_test.go b/client/internal/dns/interface_index_test.go new file mode 100644 index 000000000..9b146398a --- /dev/null +++ b/client/internal/dns/interface_index_test.go @@ -0,0 +1,35 @@ +package dns + +import ( + "net" + "testing" +) + +func TestGetInterfaceIndexExisting(t *testing.T) { + interfaces, err := net.Interfaces() + if err != nil { + t.Fatalf("list network interfaces: %v", err) + } + if len(interfaces) == 0 { + t.Fatal("expected at least one network interface") + } + + iface := interfaces[0] + index, err := getInterfaceIndex(iface.Name) + if err != nil { + t.Fatalf("look up existing interface %q: %v", iface.Name, err) + } + if index != iface.Index { + t.Fatalf("expected interface index %d, got %d", iface.Index, index) + } +} + +func TestGetInterfaceIndexMissing(t *testing.T) { + index, err := getInterfaceIndex("netbird-interface-that-does-not-exist") + if index != 0 { + t.Fatalf("expected missing interface index to be 0, got %d", index) + } + if err == nil { + t.Fatal("expected missing interface lookup to return an error") + } +} diff --git a/client/internal/dns/upstream_ios.go b/client/internal/dns/upstream_ios.go index b989bf0f9..793d87fca 100644 --- a/client/internal/dns/upstream_ios.go +++ b/client/internal/dns/upstream_ios.go @@ -130,8 +130,3 @@ func GetClientPrivate(iface privateClientIface, upstreamIP netip.Addr, dialTimeo } return client, nil } - -func getInterfaceIndex(interfaceName string) (int, error) { - iface, err := net.InterfaceByName(interfaceName) - return iface.Index, err -} From f2318a8fef230219110c9eeb58ca7f60e247ad98 Mon Sep 17 00:00:00 2001 From: evgeniyChepelev <68751844+evgeniyChepelev@users.noreply.github.com> Date: Sun, 2 Aug 2026 09:51:14 +0200 Subject: [PATCH 12/34] [client] iOS - Remove duplicate Login RPCs from the iOS SDK (#6931) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Removes redundant `Login` RPCs from the iOS SDK bindings. ## 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 an interactive iOS login option that starts authentication directly when needed. * Improved login flow handling, including clearer error reporting and successful-login notifications. * Login configuration is now saved automatically after successful authentication when applicable. * **Bug Fixes** * Prevented duplicate login requests during iOS startup. * Improved startup behavior and error propagation when the management service is unavailable. Removes redundant `Login` RPCs from the iOS SDK bindings. Both changed files are behind the `ios` build tag — Android, desktop and the shared core are not affected. ### Problem `auth.Auth.IsLoginRequired()` is not a cheap probe: it calls `doMgmLogin()` and classifies the resulting error, so every "is login required?" check costs a **full `Login` RPC**. There is no lighter way to ask. As a result the iOS client issued ~7 `Login` requests before the first `Sync`, where Android issues ~3, and the extra ones were indistinguishable from real logins in the management logs. Three of those came from this package: 1. `Run()` called `LoginSync()` before starting the engine, which performs `IsLoginRequired` **and** `Login` — two RPCs. This duplicated the engine's own `loginToManagement` (`client/internal/connect.go`), which runs immediately before the first `Sync` and is the authoritative login. The `Login(ctx, "", "")` inside `LoginSync` could not even establish anything: with an empty setup key and empty JWT, a registration attempt fails by construction, so it was a pure check. 2. `Auth.login()` called `IsLoginRequired()` again before opening the browser, even when the caller had already determined that login is needed. This is not only wasted traffic: - **It pushes peers toward the server-side login ban.** In `management/internals/shared/grpc/loginfilter.go`, every login with unchanged metadata increments `sessionCounter`, and exceeding `reconnLimitForBan` (30) within `reconnThreshold` (5 min) bans the peer for `baseBlockDuration` (10 min), doubling on repeat. Redundant logins carry identical metadata, so they count against exactly this budget. At 7 logins per connect the budget is exhausted after ~4 reconnects instead of ~10 — reachable on flaky mobile networks. - **Each redundant check is a potential 2-minute stall.** `IsLoginRequired` retries with backoff up to `MaxElapsedTime` (2 min) and returns `true` on failure, so an unreachable server was reported as "login required" rather than as a timeout, and the `LoginSync` pre-flight could abort engine startup on that basis. ### Changes **`client/ios/NetBirdSDK/client.go`** — `Run()` no longer performs the `LoginSync()` pre-flight. The engine's `loginToManagement` remains the single authoritative login. **`client/ios/NetBirdSDK/login.go`** — new exported `LoginInteractive`, which skips the `IsLoginRequired()` pre-flight and goes straight to the browser / device-code flow, for callers that have already established login is required. `LoginWithDeviceName` keeps the check for callers where the auth state is unknown (tvOS). Both now delegate to a shared `startLogin()`. ### Why this is safe An expired or revoked session still fails the connection, one step later and through a single path: `loginToManagement` returns `PermissionDenied` → the deferred `MarkManagementDisconnected` records it on the shared status recorder → `ClientStop` fires the listener's disconnect callback, where `IsLoginRequiredCached()` reports login-required → the client tears the tunnel down. The error is also returned out of `Run()`. Where the server is unreachable, the engine now retries with backoff and recovers on its own instead of aborting the start. Co-authored-by: Zoltan Papp --- client/ios/NetBirdSDK/client.go | 20 +++++++++++------ client/ios/NetBirdSDK/login.go | 38 ++++++++++++++++++++++++++------- 2 files changed, 43 insertions(+), 15 deletions(-) diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index 37d3e5d99..2d5460d03 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -158,13 +158,19 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error { defer c.ctxCancel() c.ctxCancelLock.Unlock() - auth := NewAuthWithConfig(ctx, cfg) - err = auth.LoginSync() - if err != nil { - return err - } - - log.Infof("Auth successful") + // No login pre-flight here. The engine's own loginToManagement (connect.go) performs + // the authoritative Login immediately before the first Sync, so a LoginSync() call at + // this point only duplicated it — costing two extra Login RPCs (IsLoginRequired + + // Login) on every engine start, since IsLoginRequired is itself a full Login RPC. + // + // Auth failures still reach the caller through the engine path: loginToManagement + // returns PermissionDenied, which marks the shared status recorder + // (MarkManagementDisconnected) and fires ClientStop → onDisconnected, where + // IsLoginRequiredCached() reports login-required. The error is also returned out of Run(). + // + // A pre-flight was also actively harmful when the server is unreachable: its 2-minute + // backoff blocked the start and then reported "login required" for what was really a + // timeout. The engine instead keeps retrying and recovers when the server returns. // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) c.onHostDnsFn = func([]string) {} diff --git a/client/ios/NetBirdSDK/login.go b/client/ios/NetBirdSDK/login.go index 99486839b..6cba0c411 100644 --- a/client/ios/NetBirdSDK/login.go +++ b/client/ios/NetBirdSDK/login.go @@ -222,17 +222,36 @@ func (a *Auth) Login(resultListener ErrListener, urlOpener URLOpener, forceDevic // LoginWithDeviceName performs interactive login with device authentication support // The deviceName parameter allows specifying a custom device name (required for tvOS) func (a *Auth) LoginWithDeviceName(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string) { + a.startLogin(resultListener, urlOpener, forceDeviceAuth, deviceName, false) +} + +// LoginInteractive performs the same interactive login as LoginWithDeviceName but skips the +// IsLoginRequired() pre-flight and goes straight to the browser / device-code flow. +// +// IsLoginRequired() is itself a full Login RPC against the management server, so when the +// caller has ALREADY established that login is required it is a pure duplicate. On iOS the +// main app decides to show the browser based on its own isLoginRequired() check and then +// calls straight into this method, so re-asking the server would add another Login RPC to +// every interactive login. +// +// Use LoginWithDeviceName when the auth state is unknown and a silent (browser-less) login +// must still be possible; use this when the browser is going to be shown regardless. +func (a *Auth) LoginInteractive(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string) { + a.startLogin(resultListener, urlOpener, forceDeviceAuth, deviceName, true) +} + +func (a *Auth) startLogin(resultListener ErrListener, urlOpener URLOpener, forceDeviceAuth bool, deviceName string, skipLoginCheck bool) { if resultListener == nil { - log.Errorf("LoginWithDeviceName: resultListener is nil") + log.Errorf("startLogin: resultListener is nil") return } if urlOpener == nil { - log.Errorf("LoginWithDeviceName: urlOpener is nil") + log.Errorf("startLogin: urlOpener is nil") resultListener.OnError(fmt.Errorf("urlOpener is nil")) return } go func() { - err := a.login(urlOpener, forceDeviceAuth, deviceName) + err := a.login(urlOpener, forceDeviceAuth, deviceName, skipLoginCheck) if err != nil { resultListener.OnError(err) } else { @@ -241,7 +260,7 @@ func (a *Auth) LoginWithDeviceName(resultListener ErrListener, urlOpener URLOpen }() } -func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName string) error { +func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName string, skipLoginCheck bool) error { // Create context with device name if provided ctx := a.ctx if deviceName != "" { @@ -255,10 +274,13 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin } defer authClient.Close() - // check if we need to generate JWT token - needsLogin, err := authClient.IsLoginRequired(ctx) - if err != nil { - return fmt.Errorf("failed to check login requirement: %v", err) + // check if we need to generate JWT token (skipped when the caller already knows) + needsLogin := true + if !skipLoginCheck { + needsLogin, err = authClient.IsLoginRequired(ctx) + if err != nil { + return fmt.Errorf("failed to check login requirement: %v", err) + } } jwtToken := "" From 6044663788e7455b930cefe03e26dcf1ac81ff5e Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 10:21:14 +0200 Subject: [PATCH 13/34] [client] Declare GTK4/WebKitGTK runtime deps for the Linux UI packages (#6893) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The 0.75.0 UI is built against GTK 4.14 (Ubuntu 24.04 runner) and Wails v3 calls gdk_monitor_get_scale (GTK 4.14+) unconditionally, but the deb/rpm packages only depended on netbird. On distros shipping an older GTK4 (Ubuntu 22.04, RHEL 9, openSUSE Leap 15.6) the package installed fine and then died at startup with a symbol lookup error (#6890). Declare the real runtime dependencies so package managers reject the install up front instead: - deb: libgtk-4-1 (>= 4.14) and libwebkitgtk-6.0-4 - rpm: rich (boolean) dependencies that accept both the Fedora/RHEL and the SUSE package names, with the 4.14 floor: (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) (webkitgtk6.0 or libwebkitgtk-6_0-4) Rich deps are supported by dnf and zypper (RPM 4.13+); the pinned goreleaser v2.16.0 -> nfpm v2.46.3 -> rpmpack v0.7.1 chain passes the parenthesized form through verbatim. On RHEL/Alma/Rocky 10 the webkitgtk6.0 package comes from EPEL, which becomes an install prerequisite for the UI. The wails3 packaging config (client/ui/build/linux/nfpm/nfpm.yaml, local dev packaging only) is kept in sync. ## 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 Linux package installation compatibility by requiring supported GTK 4.14 and WebKit components. * Updated Debian and RPM packages to recognize equivalent platform-specific library names. * Refined Linux package metadata to support installation across a wider range of distributions. --- .goreleaser_ui.yaml | 8 ++++++-- client/ui/build/linux/nfpm/nfpm.yaml | 8 ++++---- 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/.goreleaser_ui.yaml b/.goreleaser_ui.yaml index 1157e6379..ca5148823 100644 --- a/.goreleaser_ui.yaml +++ b/.goreleaser_ui.yaml @@ -93,7 +93,9 @@ nfpms: - src: client/ui/build/appicon.png dst: /usr/share/pixmaps/netbird.png dependencies: - - netbird + - netbird (>= 0.75.0) + - libgtk-4-1 (>= 4.14) + - libwebkitgtk-6.0-4 - maintainer: Netbird description: Netbird client UI. @@ -114,7 +116,9 @@ nfpms: - src: client/ui/build/appicon.png dst: /usr/share/pixmaps/netbird.png dependencies: - - netbird + - netbird >= 0.75.0 + - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) + - (webkitgtk6.0 or libwebkitgtk-6_0-4) rpm: signature: diff --git a/client/ui/build/linux/nfpm/nfpm.yaml b/client/ui/build/linux/nfpm/nfpm.yaml index a05daef62..764855a63 100644 --- a/client/ui/build/linux/nfpm/nfpm.yaml +++ b/client/ui/build/linux/nfpm/nfpm.yaml @@ -26,17 +26,17 @@ contents: # Default dependencies for the GTK4 + WebKitGTK 6.0 stack (Ubuntu 24.04+ / Debian 13+) depends: - - libgtk-4-1 + - libgtk-4-1 (>= 4.14) - libwebkitgtk-6.0-4 - xdg-utils # Distribution-specific overrides for different package formats overrides: - # RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux + # RPM packages for Fedora / RHEL / AlmaLinux / Rocky Linux / openSUSE rpm: depends: - - gtk4 - - webkitgtk6.0 + - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) + - (webkitgtk6.0 or libwebkitgtk-6_0-4) - xdg-utils # Arch Linux packages From 7639655883f1d3abc6dcb3440cd72152ee314fbf Mon Sep 17 00:00:00 2001 From: Brad Ison Date: Mon, 3 Aug 2026 12:26:38 +0200 Subject: [PATCH 14/34] [management] Generic gRPC extension seam for external modules (#6894) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes This adds an extension point to the management server for registering additional gRPC services. We already have a generic integrations system and dependency injection for server components. This closes the gap on being able to also extend the gRPC API cleanly. ## Issue ticket number and link N/A ## 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) No docs needed. This is strictly a small internal plumbing enhancement / refactor. --- 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 gRPC extension mechanism to contribute additional services and automatically chain extra unary and stream interceptors. * Extension shutdown hooks now run as part of server stop. * Added exported proxy token generation via `GenerateProxyToken()` for external integrations. * **Tests** * Added coverage for extension interceptor/service wiring, extension shutdown execution, and proxy token generation validation (including hash consistency and prefix). --- management/internals/server/boot.go | 9 +- management/internals/server/grpc_extension.go | 74 ++++++++ .../internals/server/grpc_extension_test.go | 160 ++++++++++++++++++ management/internals/server/server.go | 6 + management/server/types/proxy_access_token.go | 7 +- .../server/types/proxy_access_token_test.go | 17 ++ 6 files changed, 270 insertions(+), 3 deletions(-) create mode 100644 management/internals/server/grpc_extension.go create mode 100644 management/internals/server/grpc_extension_test.go diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index 1c78af9d0..2a1e521aa 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -24,13 +24,13 @@ import ( "github.com/netbirdio/netbird/encryption" "github.com/netbirdio/netbird/formatter/hook" + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/activity" activitystore "github.com/netbirdio/netbird/management/server/activity/store" - "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" nbcache "github.com/netbirdio/netbird/management/server/cache" nbContext "github.com/netbirdio/netbird/management/server/context" nbhttp "github.com/netbirdio/netbird/management/server/http" @@ -184,6 +184,10 @@ func (s *BaseServer) GRPCServer() *grpc.Server { grpc.ChainStreamInterceptor(realip.StreamServerInterceptorOpts(realipOpts...), streamInterceptor, proxyStream), } + // Append interceptors contributed by registered gRPC extensions. These + // run after the built-in chain (ChainUnaryInterceptor is additive). + gRPCOpts = appendExtensionInterceptors(gRPCOpts, s.grpcExtensions) + if s.Config.HttpConfig.LetsEncryptDomain != "" { certManager, err := encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain) if err != nil { @@ -215,6 +219,9 @@ func (s *BaseServer) GRPCServer() *grpc.Server { mgmtProto.RegisterProxyServiceServer(gRPCAPIHandler, s.ReverseProxyGRPCServer()) log.Info("ProxyService registered on gRPC server") + // Register services contributed by external modules via the extension seam. + registerExtensions(gRPCAPIHandler, s.grpcExtensions) + return gRPCAPIHandler }) } diff --git a/management/internals/server/grpc_extension.go b/management/internals/server/grpc_extension.go new file mode 100644 index 000000000..3f257c75e --- /dev/null +++ b/management/internals/server/grpc_extension.go @@ -0,0 +1,74 @@ +package server + +import ( + "context" + + "google.golang.org/grpc" +) + +// GRPCExtension bundles an external module's contribution to the management +// gRPC server: the registration of one or more services onto the shared +// grpc.Server, any server-wide interceptors those services require, and an +// optional shutdown hook. It is a generic extension point with no knowledge of +// any specific service. +type GRPCExtension struct { + // Register is invoked with the shared grpc.Server (as a ServiceRegistrar) + // after the built-in services are registered. It may register any number of + // services. May be nil. + Register func(grpc.ServiceRegistrar) + // UnaryInterceptors are appended to the server's unary interceptor chain, + // running after the built-in interceptors. May be empty. + UnaryInterceptors []grpc.UnaryServerInterceptor + // StreamInterceptors are appended to the server's stream interceptor chain, + // running after the built-in interceptors. May be empty. + StreamInterceptors []grpc.StreamServerInterceptor + // Shutdown, if non-nil, is called once during Stop() with the context + // governing server shutdown, which carries a deadline. The hook MUST + // return promptly and MUST abandon its work once that context is + // cancelled or expires: it runs before the rest of Stop()'s cleanup + // (store, event store, embedded IdP) and before Stop() itself checks the + // context's deadline, so a hook that ignores the context will delay all + // of that cleanup and prevent Stop() from returning on time. May be nil. + Shutdown func(ctx context.Context) +} + +// RegisterGRPCExtension registers a gRPC extension. Call before the gRPC server +// is first built (i.e. before Start); registrations after that have no effect. +func (s *BaseServer) RegisterGRPCExtension(ext GRPCExtension) { + s.grpcExtensions = append(s.grpcExtensions, ext) +} + +// appendExtensionInterceptors appends each extension's interceptors to the gRPC +// server options as additional chained interceptors. grpc.ChainUnaryInterceptor +// and grpc.ChainStreamInterceptor are additive, so the returned options run the +// extension interceptors after any interceptors already present in opts. +func appendExtensionInterceptors(opts []grpc.ServerOption, exts []GRPCExtension) []grpc.ServerOption { + for _, ext := range exts { + if len(ext.UnaryInterceptors) > 0 { + opts = append(opts, grpc.ChainUnaryInterceptor(ext.UnaryInterceptors...)) + } + if len(ext.StreamInterceptors) > 0 { + opts = append(opts, grpc.ChainStreamInterceptor(ext.StreamInterceptors...)) + } + } + return opts +} + +// registerExtensions registers each extension's services onto reg. +func registerExtensions(reg grpc.ServiceRegistrar, exts []GRPCExtension) { + for _, ext := range exts { + if ext.Register != nil { + ext.Register(reg) + } + } +} + +// runExtensionShutdownHooks calls each extension's shutdown hook, if set, +// passing ctx through so hooks can honor its deadline/cancellation. +func runExtensionShutdownHooks(ctx context.Context, exts []GRPCExtension) { + for _, ext := range exts { + if ext.Shutdown != nil { + ext.Shutdown(ctx) + } + } +} diff --git a/management/internals/server/grpc_extension_test.go b/management/internals/server/grpc_extension_test.go new file mode 100644 index 000000000..8f444ca72 --- /dev/null +++ b/management/internals/server/grpc_extension_test.go @@ -0,0 +1,160 @@ +package server + +import ( + "context" + "net" + "sync/atomic" + "testing" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/health" + healthgrpc "google.golang.org/grpc/health/grpc_health_v1" + "google.golang.org/grpc/test/bufconn" +) + +// Test that an extension's interceptors and service registration are actually +// wired onto a real in-process gRPC server via the helpers, and that shutdown +// hooks run. This validates the load-bearing assumption that +// grpc.ChainUnaryInterceptor is additive (extension interceptors run in +// addition to any base chain). +func TestGRPCExtensionAppliedToServer(t *testing.T) { + var unaryCalls atomic.Int32 + var streamShutdownCalled atomic.Bool + + ext := GRPCExtension{ + Register: func(reg grpc.ServiceRegistrar) { + healthgrpc.RegisterHealthServer(reg, health.NewServer()) + }, + UnaryInterceptors: []grpc.UnaryServerInterceptor{ + func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + unaryCalls.Add(1) + return handler(ctx, req) + }, + }, + Shutdown: func(ctx context.Context) { streamShutdownCalled.Store(true) }, + } + exts := []GRPCExtension{ext} + + // Base options mimic GRPCServer(): a pre-existing chain the extension appends to. + var baseUnaryCalls atomic.Int32 + opts := []grpc.ServerOption{ + grpc.ChainUnaryInterceptor(func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + baseUnaryCalls.Add(1) + return handler(ctx, req) + }), + } + opts = appendExtensionInterceptors(opts, exts) + + srv := grpc.NewServer(opts...) + registerExtensions(srv, exts) + + lis := bufconn.Listen(1024 * 1024) + go func() { _ = srv.Serve(lis) }() + t.Cleanup(srv.Stop) + + conn, err := grpc.NewClient("passthrough:///bufnet", + grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { return lis.DialContext(ctx) }), + grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = conn.Close() }) + + _, err = healthgrpc.NewHealthClient(conn).Check(context.Background(), &healthgrpc.HealthCheckRequest{}) + if err != nil { + t.Fatalf("health check via extension-registered service failed: %v", err) + } + if baseUnaryCalls.Load() != 1 { + t.Errorf("base interceptor calls = %d, want 1 (base chain must be preserved)", baseUnaryCalls.Load()) + } + if unaryCalls.Load() != 1 { + t.Errorf("extension interceptor calls = %d, want 1", unaryCalls.Load()) + } + + runExtensionShutdownHooks(context.Background(), exts) + if !streamShutdownCalled.Load() { + t.Error("extension shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookReceivesCallerContext asserts that each hook receives +// a non-nil context and that it is the very same context the caller passed +// in, so hooks can rely on values/deadlines placed on it by Stop(). +func TestGRPCExtensionShutdownHookReceivesCallerContext(t *testing.T) { + type sentinelKey struct{} + want := "shutdown-ctx-sentinel" + ctx := context.WithValue(context.Background(), sentinelKey{}, want) + + var called bool + ext := GRPCExtension{ + Shutdown: func(hookCtx context.Context) { + called = true + if hookCtx == nil { + t.Fatal("hook received a nil context") + } + got, _ := hookCtx.Value(sentinelKey{}).(string) + if got != want { + t.Errorf("hook context sentinel = %q, want %q (not the caller's context)", got, want) + } + }, + } + + runExtensionShutdownHooks(ctx, []GRPCExtension{ext}) + if !called { + t.Fatal("shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookObservesCancellation documents, by test, that +// hooks can honor cancellation/deadlines: a hook given an already-cancelled +// context must see ctx.Err() != nil and a closed Done() channel. +func TestGRPCExtensionShutdownHookObservesCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + var called bool + ext := GRPCExtension{ + Shutdown: func(hookCtx context.Context) { + called = true + if hookCtx.Err() == nil { + t.Error("hook context Err() = nil, want non-nil for a cancelled context") + } + select { + case <-hookCtx.Done(): + default: + t.Error("hook context Done() channel is not closed for a cancelled context") + } + }, + } + + runExtensionShutdownHooks(ctx, []GRPCExtension{ext}) + if !called { + t.Fatal("shutdown hook was not called") + } +} + +// TestGRPCExtensionShutdownHookNilSkipped asserts that an extension +// with a nil Shutdown hook is skipped without panicking, and that hooks for +// other extensions still run. +func TestGRPCExtensionShutdownHookNilSkipped(t *testing.T) { + var called atomic.Bool + exts := []GRPCExtension{ + {Shutdown: nil}, + {Shutdown: func(context.Context) { called.Store(true) }}, + } + + runExtensionShutdownHooks(context.Background(), exts) + if !called.Load() { + t.Error("shutdown hook for non-nil extension was not called") + } +} + +func TestRegisterGRPCExtensionAccumulates(t *testing.T) { + s := &BaseServer{} + s.RegisterGRPCExtension(GRPCExtension{}) + s.RegisterGRPCExtension(GRPCExtension{}) + if len(s.grpcExtensions) != 2 { + t.Fatalf("grpcExtensions len = %d, want 2", len(s.grpcExtensions)) + } +} diff --git a/management/internals/server/server.go b/management/internals/server/server.go index 7fd06d947..22a61bada 100644 --- a/management/internals/server/server.go +++ b/management/internals/server/server.go @@ -68,6 +68,11 @@ type BaseServer struct { proxyAuthClose func() + // grpcExtensions holds additional gRPC services, interceptors, and shutdown + // hooks registered by external modules via RegisterGRPCExtension. Populated + // during boot (single-threaded), consumed by GRPCServer() and Stop(). + grpcExtensions []GRPCExtension + listener net.Listener certManager *autocert.Manager update *version.Update @@ -257,6 +262,7 @@ func (s *BaseServer) Stop() error { s.proxyAuthClose() s.proxyAuthClose = nil } + runExtensionShutdownHooks(ctx, s.grpcExtensions) _ = s.Store().Close(ctx) _ = s.EventStore().Close(ctx) if s.update != nil { diff --git a/management/server/types/proxy_access_token.go b/management/server/types/proxy_access_token.go index b20b83bc1..9bb27ef02 100644 --- a/management/server/types/proxy_access_token.go +++ b/management/server/types/proxy_access_token.go @@ -68,7 +68,7 @@ type ProxyAccessTokenGenerated struct { // CreateNewProxyAccessToken generates a new proxy access token. // Returns the token with hashed value stored and plain token for one-time display. func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID *string, createdBy string) (*ProxyAccessTokenGenerated, error) { - hashedToken, plainToken, err := generateProxyToken() + hashedToken, plainToken, err := GenerateProxyToken() if err != nil { return nil, err } @@ -94,7 +94,10 @@ func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID * }, nil } -func generateProxyToken() (HashedProxyToken, PlainProxyToken, error) { +// GenerateProxyToken generates a new random proxy token, returning its SHA-256 +// hash (for storage) and the one-time plaintext. Exported so external modules +// can mint tokens in the canonical proxy-token format. +func GenerateProxyToken() (HashedProxyToken, PlainProxyToken, error) { secret, err := b.Random(ProxyTokenSecretLength) if err != nil { return "", "", err diff --git a/management/server/types/proxy_access_token_test.go b/management/server/types/proxy_access_token_test.go index aa1a4d2dd..740b87c2f 100644 --- a/management/server/types/proxy_access_token_test.go +++ b/management/server/types/proxy_access_token_test.go @@ -1,6 +1,7 @@ package types import ( + "strings" "testing" "time" @@ -123,6 +124,22 @@ func TestCreateNewProxyAccessToken(t *testing.T) { }) } +func TestGenerateProxyToken(t *testing.T) { + hashed, plain, err := GenerateProxyToken() + if err != nil { + t.Fatal(err) + } + if err := plain.Validate(); err != nil { + t.Errorf("generated token failed Validate(): %v", err) + } + if plain.Hash() != hashed { + t.Error("returned hashed token does not match Hash(plain)") + } + if !strings.HasPrefix(string(plain), ProxyTokenPrefix) { + t.Errorf("token %q missing prefix %q", plain, ProxyTokenPrefix) + } +} + func TestProxyAccessToken_IsExpired(t *testing.T) { past := time.Now().Add(-1 * time.Hour) future := time.Now().Add(1 * time.Hour) From 2bfd9fcffe40acdf3d7ae7bf285f7f14a3931fb3 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Mon, 3 Aug 2026 19:45:17 +0900 Subject: [PATCH 15/34] [management] Resolve agent network permissions per submodule (#7030) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Agent Network gates providers, policies, guardrails, budgets, usage, access logs, and settings behind the single `agent_network` permission module, so access is all-or-nothing: a future delegated role cannot be scoped to a subset of the area (for example usage-only visibility). This introduces dotted submodules (`agent_network.providers`, `.policies`, `.guardrails`, `.budgets`, `.usage`, `.logs`, `.settings`) and resolves grants with a cascade: exact module first, then its parent, then the role's `AutoAllowNew` default. The agent network manager now validates each operation against its matching submodule. `usage` (aggregated counters, overview) is deliberately separate from `logs` (request-level entries, which can contain captured prompts). No role definitions change. No built-in role carries an explicit `agent_network` entry, so every role resolves the submodules exactly as it resolved the parent module before — pinned by a test that compares each built-in role's answer on every submodule against its answer on `agent_network`. Role additions that use these submodules come separately. --- .../internals/modules/agentnetwork/manager.go | 77 ++++++---- .../agentnetwork/provider_bootstrap_test.go | 134 +++++++++++++++++ management/server/permissions/manager.go | 22 ++- management/server/permissions/manager_test.go | 139 ++++++++++++++++++ .../server/permissions/modules/module.go | 30 ++++ 5 files changed, 372 insertions(+), 30 deletions(-) create mode 100644 management/internals/modules/agentnetwork/provider_bootstrap_test.go create mode 100644 management/server/permissions/manager_test.go diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index 77c77ce44..2687e4534 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -157,14 +157,14 @@ func NewManager( } func (m *managerImpl) GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID) @@ -175,9 +175,14 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid // been created yet; otherwise it is ignored (the cluster is pinned on // Settings and every provider in the account routes through it). func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error) { - if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Create); err != nil { return nil, err } + if strings.TrimSpace(bootstrapCluster) != "" { + if err := m.requireSettingsBootstrapPermission(ctx, provider.AccountID, userID); err != nil { + return nil, err + } + } // An empty api_key would silently produce a synthesised service // that 401s on every upstream request. Surface the misconfiguration @@ -218,7 +223,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide } func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error) { - if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Update); err != nil { return nil, err } @@ -257,7 +262,7 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide } func (m *managerImpl) DeleteProvider(ctx context.Context, accountID, userID, providerID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Delete); err != nil { return err } @@ -306,21 +311,21 @@ func pluralize(n int, singular, plural string) string { } func (m *managerImpl) GetAllPolicies(ctx context.Context, accountID, userID string) ([]*types.Policy, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetPolicy(ctx context.Context, accountID, userID, policyID string) (*types.Policy, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkPolicyByID(ctx, store.LockingStrengthNone, accountID, policyID) } func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) { - if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Create); err != nil { return nil, err } @@ -346,7 +351,7 @@ func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *t } func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) { - if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Update); err != nil { return nil, err } @@ -373,7 +378,7 @@ func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *t } func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, policyID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Delete); err != nil { return err } @@ -393,21 +398,21 @@ func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, polic } func (m *managerImpl) GetAllGuardrails(ctx context.Context, accountID, userID string) ([]*types.Guardrail, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID) } func (m *managerImpl) GetGuardrail(ctx context.Context, accountID, userID, guardrailID string) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkGuardrailByID(ctx, store.LockingStrengthNone, accountID, guardrailID) } func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Create); err != nil { return nil, err } @@ -429,7 +434,7 @@ func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardr } func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) { - if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Update); err != nil { return nil, err } @@ -452,7 +457,7 @@ func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardr } func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, guardrailID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Delete); err != nil { return err } @@ -473,7 +478,7 @@ func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, gu // GetAllBudgetRules returns every account-level budget rule for the account. func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID string) ([]*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil { return nil, err } return m.store.GetAccountAgentNetworkBudgetRules(ctx, store.LockingStrengthNone, accountID) @@ -481,7 +486,7 @@ func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID s // GetBudgetRule returns a single account-level budget rule. func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, ruleID string) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkBudgetRuleByID(ctx, store.LockingStrengthNone, accountID, ruleID) @@ -491,7 +496,7 @@ func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, rule // enforced at request time (CheckLLMPolicyLimits), not baked into the synth // proxy config, so no reconcile is needed. func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Create); err != nil { + if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Create); err != nil { return nil, err } @@ -513,7 +518,7 @@ func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule // UpdateBudgetRule updates an existing account-level budget rule. func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) { - if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Update); err != nil { return nil, err } @@ -536,7 +541,7 @@ func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule // DeleteBudgetRule removes an account-level budget rule. func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, ruleID string) error { - if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Delete); err != nil { return err } @@ -561,7 +566,7 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r // gating, access-log emission), a reconcile is triggered so the proxy and peer // network maps converge on the new state. func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) { - if err := m.requirePermission(ctx, settings.AccountID, userID, operations.Update); err != nil { + if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil { return nil, err } @@ -615,7 +620,7 @@ func (m *managerImpl) validateProviderRefs(ctx context.Context, accountID string // Returns the underlying status.NotFound when no row has been // bootstrapped yet (i.e. the account has no providers). func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Read); err != nil { return nil, err } return m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) @@ -627,6 +632,22 @@ func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) // the subdomain is picked from the curated wordlist avoiding // collisions on the same cluster. Idempotent: if a row already exists // it is returned untouched and the hint is ignored. +// requireSettingsBootstrapPermission gates the one-time settings bootstrap a +// first provider create performs. Pinning the account's cluster and subdomain +// is a settings write, so it needs the settings permission on top of the +// provider one. No-op once the settings row exists. +func (m *managerImpl) requireSettingsBootstrapPermission(ctx context.Context, accountID, userID string) error { + _, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + if err == nil { + return nil + } + var sErr *status.Error + if !errors.As(err, &sErr) || sErr.Type() != status.NotFound { + return fmt.Errorf("get agent network settings: %w", err) + } + return m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Create) +} + func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, providerCluster string) (*types.Settings, error) { if accountID == "" { return nil, fmt.Errorf("bootstrap settings: account id is required") @@ -685,7 +706,7 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, // counter view; permission gate is the same Read role that gates // every other agent-network surface. func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID string) ([]*types.Consumption, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil { return nil, err } return m.store.ListAgentNetworkConsumption(ctx, store.LockingStrengthNone, accountID) @@ -694,7 +715,7 @@ func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID str // ListAccessLogs returns a paginated, server-side-filtered page of // agent-network access logs plus the total count matching the filter. func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLog, int64, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil { return nil, 0, err } return m.store.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID, filter) @@ -704,7 +725,7 @@ func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID stri // agent-network access logs grouped by session, plus the total number of // sessions matching the filter. func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil { return nil, 0, err } return m.store.GetAgentNetworkAccessLogSessions(ctx, store.LockingStrengthNone, accountID, filter) @@ -713,7 +734,7 @@ func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, user // GetUsageOverview returns the filtered usage rows aggregated into time buckets // at the requested granularity, oldest-first. func (m *managerImpl) GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error) { - if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil { + if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil { return nil, err } rows, err := m.store.GetAgentNetworkUsageRows(ctx, store.LockingStrengthNone, accountID, filter) @@ -787,8 +808,8 @@ func (m *managerImpl) RecordConsumption(ctx context.Context, accountID string, k return m.store.IncrementAgentNetworkConsumption(ctx, accountID, kind, dimID, windowSeconds, windowStart, tokensIn, tokensOut, costUSD) } -func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, op operations.Operation) error { - ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetwork, op) +func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, module modules.Module, op operations.Operation) error { + ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, module, op) if err != nil { return status.NewPermissionValidationError(err) } diff --git a/management/internals/modules/agentnetwork/provider_bootstrap_test.go b/management/internals/modules/agentnetwork/provider_bootstrap_test.go new file mode 100644 index 000000000..1a2904c51 --- /dev/null +++ b/management/internals/modules/agentnetwork/provider_bootstrap_test.go @@ -0,0 +1,134 @@ +package agentnetwork + +import ( + "context" + "runtime" + "testing" + + "github.com/golang/mock/gomock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/account" + "github.com/netbirdio/netbird/management/server/permissions" + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/store" + nbtypes "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +// bootstrapFixture wires a real sqlite store to a gomock permissions manager +// so tests can grant the provider permission while denying (or never +// expecting) the settings one. +type bootstrapFixture struct { + manager Manager + store store.Store + perms *permissions.MockManager +} + +func newBootstrapFixture(t *testing.T) *bootstrapFixture { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("sqlite store not properly supported on Windows yet") + } + t.Setenv("NETBIRD_STORE_ENGINE", string(nbtypes.SqliteStoreEngine)) + + st, cleanUp, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) + require.NoError(t, err, "test store setup must succeed") + t.Cleanup(cleanUp) + + ctrl := gomock.NewController(t) + perms := permissions.NewMockManager(ctrl) + + accounts := account.NewMockManager(ctrl) + accounts.EXPECT().StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + accounts.EXPECT().UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + accounts.EXPECT().BufferUpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes() + + return &bootstrapFixture{ + manager: NewManager(st, perms, accounts, nil), + store: st, + perms: perms, + } +} + +func (f *bootstrapFixture) expectPermission(accountID, userID string, module modules.Module, op operations.Operation, allowed bool) { + f.perms.EXPECT(). + ValidateUserPermissions(gomock.Any(), accountID, userID, module, op). + Return(allowed, context.Background(), nil) +} + +func newBootstrapProvider(accountID string) *types.Provider { + p := types.NewProvider(accountID) + p.Name = "openai" + p.UpstreamURL = "https://api.openai.com" + p.APIKey = "sk-test" + p.Enabled = true + return p +} + +// TestCreateProviderBootstrapRequiresSettingsPermission pins the gate on the +// one-time settings bootstrap: creating the first provider with a +// bootstrap_cluster pins the account's cluster and subdomain, which is a +// settings write and must not ride on the providers permission alone. +func TestCreateProviderBootstrapRequiresSettingsPermission(t *testing.T) { + ctx := context.Background() + + t.Run("denied without settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, false) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.Error(t, err, "bootstrap without settings permission must fail") + var sErr *status.Error + require.ErrorAs(t, err, &sErr) + assert.Equal(t, status.PermissionDenied, sErr.Type(), "denial should surface as permission denied") + + providers, err := f.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "account1") + require.NoError(t, err) + assert.Empty(t, providers, "provider must not be persisted when bootstrap is denied") + _, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1") + assert.Error(t, err, "settings row must not be created when bootstrap is denied") + }) + + t.Run("allowed with settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true) + + created, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.NoError(t, err, "bootstrap with both permissions must succeed") + require.NotNil(t, created) + + settings, err := f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1") + require.NoError(t, err, "bootstrap must create the settings row") + assert.Equal(t, "cluster1.example.com", settings.Cluster, "settings should pin the bootstrap cluster") + }) + + t.Run("existing settings need no settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + require.NoError(t, f.store.SaveAgentNetworkSettings(ctx, &types.Settings{ + AccountID: "account1", + Cluster: "cluster1.example.com", + Subdomain: "existing", + }), "pre-existing settings row setup must succeed") + + // Only the providers permission may be consulted: gomock fails the + // test on any unexpected settings-permission call. + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com") + require.NoError(t, err, "create with existing settings must not require the settings permission") + }) + + t.Run("no bootstrap cluster needs no settings permission", func(t *testing.T) { + f := newBootstrapFixture(t) + f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true) + + _, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "") + require.NoError(t, err, "create without bootstrap must not require the settings permission") + }) +} diff --git a/management/server/permissions/manager.go b/management/server/permissions/manager.go index 995f234d8..6b9977a86 100644 --- a/management/server/permissions/manager.go +++ b/management/server/permissions/manager.go @@ -82,6 +82,9 @@ func (m *managerImpl) ValidateUserPermissions( return m.ValidateRoleModuleAccess(ctx, accountID, role, module, operation), ctxEnriched, nil } +// ValidateRoleModuleAccess resolves an operation against the role's explicit +// grant for the module, then the grant for its parent module when the module +// is a dotted submodule, and finally the role's AutoAllowNew default. func (m *managerImpl) ValidateRoleModuleAccess( ctx context.Context, accountID string, @@ -89,7 +92,7 @@ func (m *managerImpl) ValidateRoleModuleAccess( module modules.Module, operation operations.Operation, ) bool { - if permissions, ok := role.Permissions[module]; ok { + if permissions, ok := lookupModulePermissions(role, module); ok { if allowed, exists := permissions[operation]; exists { return allowed } @@ -100,6 +103,21 @@ func (m *managerImpl) ValidateRoleModuleAccess( return role.AutoAllowNew[operation] } +// lookupModulePermissions returns the role's explicit permission set for the +// module, falling back to the parent module's set for dotted submodules. The +// second return reports whether any explicit set was found. +func lookupModulePermissions(role roles.RolePermissions, module modules.Module) (map[operations.Operation]bool, bool) { + if permissions, ok := role.Permissions[module]; ok { + return permissions, true + } + if parent, hasParent := module.Parent(); hasParent { + if permissions, ok := role.Permissions[parent]; ok { + return permissions, true + } + } + return nil, false +} + func (m *managerImpl) ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) (context.Context, error) { if user.AccountID != accountID { return ctx, status.NewUserNotPartOfAccountError() @@ -119,7 +137,7 @@ func (m *managerImpl) GetPermissionsByRole(ctx context.Context, role types.UserR permissions := roles.Permissions{} for k := range modules.All { - if rolePermissions, ok := roleMap.Permissions[k]; ok { + if rolePermissions, ok := lookupModulePermissions(roleMap, k); ok { permissions[k] = rolePermissions continue } diff --git a/management/server/permissions/manager_test.go b/management/server/permissions/manager_test.go new file mode 100644 index 000000000..345212f43 --- /dev/null +++ b/management/server/permissions/manager_test.go @@ -0,0 +1,139 @@ +package permissions + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/permissions/roles" + "github.com/netbirdio/netbird/management/server/types" +) + +func TestValidateRoleModuleAccessSubmoduleCascade(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + fullAccess := map[operations.Operation]bool{ + operations.Read: true, + operations.Create: true, + operations.Update: true, + operations.Delete: true, + } + readOnly := map[operations.Operation]bool{ + operations.Read: true, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + } + denyAll := map[operations.Operation]bool{ + operations.Read: false, + operations.Create: false, + operations.Update: false, + operations.Delete: false, + } + + t.Run("parent grant covers submodules", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{modules.AgentNetwork: fullAccess}, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Create), + "parent full grant should allow create on a submodule") + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read), + "parent full grant should allow read on a submodule") + }) + + t.Run("submodule grant does not leak to parent or siblings", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{modules.AgentNetworkUsage: readOnly}, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read), + "explicit submodule read should be allowed") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Create), + "read-only submodule grant should not allow create") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetwork, operations.Read), + "submodule grant should not grant the parent module") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read), + "submodule grant should not grant a sibling submodule") + }) + + t.Run("explicit submodule entry wins over parent grant", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: denyAll, + Permissions: roles.Permissions{ + modules.AgentNetwork: fullAccess, + modules.AgentNetworkLogs: denyAll, + }, + } + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read), + "explicit submodule deny should override the parent grant") + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read), + "sibling submodules should still resolve through the parent grant") + }) + + t.Run("auto allow applies when neither submodule nor parent is granted", func(t *testing.T) { + role := roles.RolePermissions{ + AutoAllowNew: readOnly, + } + assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read), + "auto-allow read should apply to submodules") + assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Delete), + "auto-allow should not grant unlisted operations") + }) +} + +// TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules pins the behavior the +// submodule split must not change: every built-in role resolves the new +// submodules exactly as it resolved the agent_network module before. +func TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + submodules := []modules.Module{ + modules.AgentNetworkProviders, + modules.AgentNetworkPolicies, + modules.AgentNetworkGuardrails, + modules.AgentNetworkBudgets, + modules.AgentNetworkUsage, + modules.AgentNetworkLogs, + modules.AgentNetworkSettings, + } + allOperations := []operations.Operation{operations.Read, operations.Create, operations.Update, operations.Delete} + + for _, role := range []types.UserRole{types.UserRoleOwner, types.UserRoleAdmin, types.UserRoleAuditor, types.UserRoleNetworkAdmin, types.UserRoleUser} { + rolePermissions, ok := roles.RolesMap[role] + require.True(t, ok, "role %s must exist in RolesMap", role) + + for _, sub := range submodules { + for _, op := range allOperations { + expected := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, modules.AgentNetwork, op) + actual := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, sub, op) + assert.Equal(t, expected, actual, "role %s: %s on %s should match the agent_network module", role, op, sub) + } + } + } +} + +func TestGetPermissionsByRoleIncludesSubmodules(t *testing.T) { + manager := NewManager(nil) + ctx := context.Background() + + permissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAuditor) + require.NoError(t, err, "auditor role must resolve") + + usage, ok := permissions[modules.AgentNetworkUsage] + require.True(t, ok, "permissions map should contain the usage submodule") + assert.True(t, usage[operations.Read], "auditor should read the usage submodule") + assert.False(t, usage[operations.Update], "auditor should not update the usage submodule") + + adminPermissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAdmin) + require.NoError(t, err, "admin role must resolve") + providers, ok := adminPermissions[modules.AgentNetworkProviders] + require.True(t, ok, "permissions map should contain the providers submodule") + assert.True(t, providers[operations.Delete], "admin should delete on the providers submodule") +} diff --git a/management/server/permissions/modules/module.go b/management/server/permissions/modules/module.go index a3a9c554d..8a2a1a52d 100644 --- a/management/server/permissions/modules/module.go +++ b/management/server/permissions/modules/module.go @@ -1,5 +1,7 @@ package modules +import "strings" + type Module string const ( @@ -20,6 +22,17 @@ const ( IdentityProviders Module = "identity_providers" Services Module = "services" AgentNetwork Module = "agent_network" + + // Agent Network submodules. A role may grant one of these directly + // or grant the AgentNetwork parent, which covers all of them (see + // permissions.Manager cascade resolution). + AgentNetworkProviders Module = "agent_network.providers" + AgentNetworkPolicies Module = "agent_network.policies" + AgentNetworkGuardrails Module = "agent_network.guardrails" + AgentNetworkBudgets Module = "agent_network.budgets" + AgentNetworkUsage Module = "agent_network.usage" + AgentNetworkLogs Module = "agent_network.logs" + AgentNetworkSettings Module = "agent_network.settings" ) var All = map[Module]struct{}{ @@ -40,4 +53,21 @@ var All = map[Module]struct{}{ IdentityProviders: {}, Services: {}, AgentNetwork: {}, + + AgentNetworkProviders: {}, + AgentNetworkPolicies: {}, + AgentNetworkGuardrails: {}, + AgentNetworkBudgets: {}, + AgentNetworkUsage: {}, + AgentNetworkLogs: {}, + AgentNetworkSettings: {}, +} + +// Parent returns the module owning a dotted submodule name and true, or the +// module itself and false when it has no parent. +func (m Module) Parent() (Module, bool) { + if i := strings.IndexByte(string(m), '.'); i > 0 { + return Module(string(m)[:i]), true + } + return m, false } From f9b412228e19b525123e10d879cadb08dfeec645 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Mon, 3 Aug 2026 13:20:53 +0200 Subject: [PATCH 16/34] [management] fix handling of empty network map during decode and encode (#6987) --- .../shared/grpc/components_encoder.go | 2 + .../shared/grpc/components_encoder_test.go | 11 +++- shared/management/networkmap/decode.go | 14 +++-- shared/management/networkmap/envelope_test.go | 60 +++++++++++++++++++ 4 files changed, 80 insertions(+), 7 deletions(-) diff --git a/management/internals/shared/grpc/components_encoder.go b/management/internals/shared/grpc/components_encoder.go index 7e43cf478..e1a5cae48 100644 --- a/management/internals/shared/grpc/components_encoder.go +++ b/management/internals/shared/grpc/components_encoder.go @@ -61,6 +61,8 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel return &proto.NetworkMapEnvelope{ Payload: &proto.NetworkMapEnvelope_Full{ Full: &proto.NetworkMapComponentsFull{ + Serial: networkSerial(c.Network), + Network: toAccountNetwork(c.Network), PeerConfig: in.PeerConfig, // components.Peers always contains the target peer Peers: []*proto.PeerCompact{toPeerCompact(c.Peers[c.PeerID])}, diff --git a/management/internals/shared/grpc/components_encoder_test.go b/management/internals/shared/grpc/components_encoder_test.go index 100ab0948..f7df82f2f 100644 --- a/management/internals/shared/grpc/components_encoder_test.go +++ b/management/internals/shared/grpc/components_encoder_test.go @@ -758,6 +758,9 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) { assert.Equal(t, "netbird.cloud", full.DnsDomain) assert.Len(t, full.Peers, 1) assert.Empty(t, full.Policies) + require.NotNil(t, full.Network, "client runs Calculate() over the envelope and dereferences Network unconditionally; a nil here would crash the receiver") + assert.Equal(t, "net-empty", full.Network.Identifier) + assert.Equal(t, uint64(9), full.Serial) } func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) { @@ -776,6 +779,12 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) { func emptyNetworkMapComponents() *types.NetworkMapComponents { return types.EmptyNetworkMapComponents( &types.NetworkMapComponents{ - PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}}, + PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}, + Network: &types.Network{ + Identifier: "net-empty", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 9, + }, + }, ) } diff --git a/shared/management/networkmap/decode.go b/shared/management/networkmap/decode.go index d15117b6e..4864a9dff 100644 --- a/shared/management/networkmap/decode.go +++ b/shared/management/networkmap/decode.go @@ -228,15 +228,17 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents, return c, nil } +// decodeAccountNetwork never returns nil — Calculate() dereferences +// c.Network unconditionally, and servers that predate the fix omit the field +// entirely from the empty-components envelope. func decodeAccountNetwork(an *proto.AccountNetwork) *types.Network { + n := &types.Network{} if an == nil { - return nil - } - n := &types.Network{ - Identifier: an.Identifier, - Dns: an.Dns, - Serial: an.Serial, + return n } + n.Identifier = an.Identifier + n.Dns = an.Dns + n.Serial = an.Serial if an.NetCidr != "" { if _, ipnet, err := net.ParseCIDR(an.NetCidr); err == nil && ipnet != nil { n.Net = *ipnet diff --git a/shared/management/networkmap/envelope_test.go b/shared/management/networkmap/envelope_test.go index a81478aff..92e1916da 100644 --- a/shared/management/networkmap/envelope_test.go +++ b/shared/management/networkmap/envelope_test.go @@ -221,6 +221,66 @@ func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) { "client-side Calculate must connect the same remote peers as the server") } +// TestEnvelopeToNetworkMap_EmptyComponents covers the graceful-degrade path +// the server takes for a peer that is missing from the account or absent from +// the validated-peers map. The legacy server short-circuited before +// Calculate() and shipped a NetworkMap carrying only the account Network; the +// components path runs Calculate() on the client instead, so the envelope must +// carry Network or the client panics dereferencing a nil *types.Network. +func TestEnvelopeToNetworkMap_EmptyComponents(t *testing.T) { + localPeerKey := randomWgKey(t) + c := types.EmptyNetworkMapComponents(&types.NetworkMapComponents{ + PeerID: "peer-A", + Network: &types.Network{ + Identifier: "net-empty", + Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}, + Serial: 7, + }, + Peers: map[string]*types.ComponentPeer{ + "peer-A": {ID: "peer-A", Key: localPeerKey, IP: netip.AddrFrom4([4]byte{100, 64, 0, 1})}, + }, + }) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + require.NotNil(t, envelope.GetFull().Network, "empty envelope must carry the account Network") + + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + require.NoError(t, err, "EnvelopeToNetworkMap must degrade gracefully on empty components") + require.Equal(t, uint64(7), result.NetworkMap.Serial) + require.Empty(t, result.NetworkMap.RemotePeers, "unvalidated peer connects to nobody") +} + +// TestEnvelopeToNetworkMap_MissingNetwork simulates a server that omits +// AccountNetwork from the envelope. Clients must degrade rather than panic, so +// they survive talking to a management server that predates the encoder fix. +func TestEnvelopeToNetworkMap_MissingNetwork(t *testing.T) { + c, localPeerKey := buildSmokeComponents(t) + + envelope := mgmtgrpc.EncodeNetworkMapEnvelope(mgmtgrpc.ComponentsEnvelopeInput{ + Components: c, + DNSDomain: "netbird.cloud", + }) + envelope.GetFull().Network = nil + + wire, err := goproto.Marshal(envelope) + require.NoError(t, err, "marshal envelope") + var decoded proto.NetworkMapEnvelope + require.NoError(t, goproto.Unmarshal(wire, &decoded), "unmarshal envelope") + + result, err := nbnetworkmap.EnvelopeToNetworkMap(context.Background(), &decoded, localPeerKey, "netbird.cloud") + require.NoError(t, err, "a missing AccountNetwork must not panic the client") + require.NotNil(t, result.Components.Network) + require.NotEmpty(t, result.NetworkMap.RemotePeers, "the rest of the snapshot stays usable") +} + // buildSmokeComponents returns a minimal NetworkMapComponents (2 peers, 1 // group, 1 allow policy) plus the receiving peer's WG public key. Sufficient // to validate the encode → marshal → decode → Calculate pipeline produces From 2f721ec0d5d213e59d33aeedf1a9fd90370d2535 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 14:21:40 +0200 Subject: [PATCH 17/34] [client] Keep the account email backing the SSO login hint correct (#6986) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Three fixes to how the desktop GUI keeps the account email that backs the SSO `login_hint`. Each is independent and reviewable on its own. **1. Store the email after a GUI SSO login** The daemon returns the authenticated user's email from `WaitSSOLogin` but cannot persist it: it runs as root while the per-profile state file is user-owned. The CLI's `handleSSOLogin` writes it after its own `WaitSSOLogin`; the GUI path read the value and dropped it. The profile was therefore left with no email, so `Profiles.List` showed no account for it, and later logins and session extends went out with no `login_hint` — leaving the IdP to pick an account instead of reusing the one the profile belongs to. Mirror the CLI and store it, next to the `Logout` path that already clears the same file for the same reason. **2. File the email against the profile the login ran for** `SetActiveProfileState` resolves the target itself, so it writes to whichever profile is active when it is called. A GUI SSO login spans seconds of user interaction in the browser, and the tray stays clickable throughout: switching profiles in that window left the email filed under the profile that happened to be active when the flow returned. The wrong profile then advertised an account it does not own, and offered it as the `login_hint` next time. Adds `SetProfileState(id, state)`, the write-side counterpart of the existing `GetProfileState(id)`, and keeps `SetActiveProfileState` as a wrapper for callers with no particular profile in mind. `Login` now reports the profile it resolved so the frontend can hand it back with the SSO wait, which closes the window. **3. Delete the email when a profile is removed** Removing a profile left its state file behind: the daemon deletes what it owns, but the email file is user-owned and out of reach for a root daemon — the same split that already puts the `Logout` cleanup on the UI side. Beyond the stray file, legacy profiles are keyed by name rather than by a generated ID, so recreating a profile under a removed one's name inherited its email — shown as the account in the profile list and sent as the `login_hint` on the next login. ## 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/__ ## Summary by CodeRabbit - **New Features** - SSO login details can now be saved to the specific profile selected during sign-in. - Profile state can be managed independently for different profiles. - **Bug Fixes** - Removing a profile now also cleans up its associated saved state. - Cleanup issues no longer prevent successful profile removal and are handled gracefully. - Existing active-profile behavior remains unchanged. --- client/internal/profilemanager/state.go | 38 ++++++++++++++-------- client/ui/frontend/src/lib/connection.ts | 9 ++++-- client/ui/services/connection.go | 40 ++++++++++++++++++++++++ client/ui/services/profile.go | 26 +++++++++++++-- 4 files changed, 96 insertions(+), 17 deletions(-) diff --git a/client/internal/profilemanager/state.go b/client/internal/profilemanager/state.go index 9e9577796..fcd1c384c 100644 --- a/client/internal/profilemanager/state.go +++ b/client/internal/profilemanager/state.go @@ -45,12 +45,35 @@ func (pm *ProfileManager) GetProfileState(id ID) (*ProfileState, error) { return &state, nil } -func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { +// SetProfileState writes the state file of the profile identified by id. Prefer +// it over SetActiveProfileState whenever the caller knows which profile the data +// belongs to: an SSO login spans seconds of user interaction, and the active +// profile can change during it, which would file the account email under +// whichever profile happened to be active when the flow returned. +func (pm *ProfileManager) SetProfileState(id ID, state *ProfileState) error { configDir, err := getConfigDir() if err != nil { return fmt.Errorf("get config directory: %w", err) } + if id == "" { + return fmt.Errorf("empty profile ID") + } + if id != defaultProfileName && !IsValidProfileFilenameStem(id) { + return fmt.Errorf("invalid profile ID: %q", id) + } + + stateFile := filepath.Join(configDir, id.String()+".state.json") + if err := util.WriteJsonWithRestrictedPermission(context.Background(), stateFile, state); err != nil { + return fmt.Errorf("write profile state: %w", err) + } + + return nil +} + +// SetActiveProfileState writes the state file of whichever profile is active at +// call time. Use SetProfileState when the target profile is known. +func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { activeProf, err := pm.GetActiveProfile() if err != nil { if errors.Is(err, ErrNoActiveProfile) { @@ -59,18 +82,7 @@ func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error { return fmt.Errorf("get active profile: %w", err) } - id := activeProf.ID - if id != defaultProfileName && !IsValidProfileFilenameStem(id) { - return fmt.Errorf("invalid active profile ID: %q", id) - } - - stateFile := filepath.Join(configDir, id.String()+".state.json") - err = util.WriteJsonWithRestrictedPermission(context.Background(), stateFile, state) - if err != nil { - return fmt.Errorf("write profile state: %w", err) - } - - return nil + return pm.SetProfileState(activeProf.ID, state) } // RemoveProfileState deletes the per-profile state file (which holds the diff --git a/client/ui/frontend/src/lib/connection.ts b/client/ui/frontend/src/lib/connection.ts index fca03fc87..cc0e67cb3 100644 --- a/client/ui/frontend/src/lib/connection.ts +++ b/client/ui/frontend/src/lib/connection.ts @@ -43,7 +43,12 @@ function buildSsoCancelPromise(state: SsoState, signal?: AbortSignal): Promise { @@ -56,7 +61,7 @@ async function runSsoLogin( // suspended, so a frontend-driven Up (a promise continuation) would not // fire until the user woke the window (e.g. hovering the tray icon). const waitPromise = Connection.WaitSSOLoginAndUp( - { userCode: result.userCode, hostname: "" }, + { userCode: result.userCode, hostname: "", profileId: result.profileId }, { profileName: "", username: "" }, ); diff --git a/client/ui/services/connection.go b/client/ui/services/connection.go index fae7ddd23..1069f8754 100644 --- a/client/ui/services/connection.go +++ b/client/ui/services/connection.go @@ -33,12 +33,21 @@ type LoginResult struct { UserCode string `json:"userCode"` VerificationURI string `json:"verificationUri"` VerificationURIComplete string `json:"verificationUriComplete"` + // ProfileID is the ID of the profile this login ran against, or "" when the + // caller named the profile itself and no ID was resolved. Pass it back in + // WaitSSOParams so the account email lands on this profile even if the + // active one changes during SSO. + ProfileID string `json:"profileId"` } // WaitSSOParams are the inputs to waitSSOLogin. type WaitSSOParams struct { UserCode string `json:"userCode"` Hostname string `json:"hostname"` + // ProfileID is the profile the login was started for, used to file the + // account email against it rather than against whichever profile is active + // when the flow returns. Optional: empty falls back to the active profile. + ProfileID string `json:"profileId"` } // UpParams selects the profile to bring up. @@ -77,11 +86,16 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err // Fall back to the daemon's active profile and the current OS user. profileName := p.ProfileName username := p.Username + // Only set when the daemon told us the ID. A caller-supplied ProfileName is + // a handle — a display name or an ID prefix resolve too — and the state file + // is named after the ID, so passing a handle on would name the wrong file. + profileID := "" if profileName == "" { if active, aerr := cli.GetActiveProfile(ctx, &proto.GetActiveProfileRequest{}); aerr == nil { // Address the active profile by ID (the daemon resolves it as a // handle); names can collide, the ID cannot. profileName = active.GetId() + profileID = profileName if username == "" { username = active.GetUsername() } @@ -122,6 +136,7 @@ func (s *Connection) Login(ctx context.Context, p LoginParams) (LoginResult, err UserCode: resp.GetUserCode(), VerificationURI: resp.GetVerificationURI(), VerificationURIComplete: resp.GetVerificationURIComplete(), + ProfileID: profileID, }, nil } @@ -242,6 +257,31 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string, return "", s.classifyDaemonError(err) } log.Infof("SSO login completed, daemon reported success") + + // Persist the account email the same way the CLI does after its own + // WaitSSOLogin: the daemon returns it but cannot store it, since it runs as + // root and the per-profile state file is user-owned (see Logout below). + // Without this the profile has no email, so Profiles.List shows no account + // and later logins and session extends go out without a login_hint — + // leaving the IdP to guess which account was meant. + if email := resp.GetEmail(); email != "" { + state := &profilemanager.ProfileState{Email: email} + pm := profilemanager.NewProfileManager() + + // Against the profile the login was started for: SSO spans seconds of + // user interaction, and a profile switch in that window would otherwise + // file the email under the wrong profile. + if p.ProfileID != "" { + err = pm.SetProfileState(profilemanager.ID(p.ProfileID), state) + } else { + err = pm.SetActiveProfileState(state) + } + if err != nil { + // Non-fatal: the login itself succeeded. + log.Warnf("failed to store account email: %v", err) + } + } + return resp.GetEmail(), nil } diff --git a/client/ui/services/profile.go b/client/ui/services/profile.go index 09468c9df..5a9a0e68d 100644 --- a/client/ui/services/profile.go +++ b/client/ui/services/profile.go @@ -6,6 +6,8 @@ import ( "context" "os/user" + log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/proto" ) @@ -151,11 +153,31 @@ func (s *Profiles) Remove(ctx context.Context, p ProfileRef) error { if err != nil { return err } - _, err = cli.RemoveProfile(ctx, &proto.RemoveProfileRequest{ + resp, err := cli.RemoveProfile(ctx, &proto.RemoveProfileRequest{ ProfileName: p.ProfileName, Username: p.Username, }) - return err + if err != nil { + return err + } + + // The daemon deletes what it owns but runs as root, so it leaves the + // user-owned state file holding the account email behind (same split as + // Connection.Logout). Legacy profiles are keyed by name rather than by a + // generated ID, so a recreated profile of the same name would inherit the + // deleted one's email and offer it as the login_hint. + // + // Keyed on the ID the daemon resolved, not on the request handle: that may + // have been a display name or an ID prefix, which would name a different + // file (or none). + if id := resp.GetId(); id != "" { + if err := profilemanager.NewProfileManager().RemoveProfileState(id); err != nil { + // Non-fatal: the profile itself is gone. + log.Warnf("failed to remove profile state for %s: %v", id, err) + } + } + + return nil } // Rename changes a profile's display name. The on-disk ID is unaffected, so From 28197e6504418622436425e584d43eedeeb8ce11 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 14:30:41 +0200 Subject: [PATCH 18/34] [client] Android - Create the Android fake IP manager lazily on DNS flag enable (#6989) The fake IP manager was only created at route manager construction, from the DNS feature flag fetched by the initial GetNetworkMap call. When the flag flipped to true mid-session, UpdateRoutes set useNewDNSRoute but never created the manager, so domain routes added after the flip got a DNS interceptor with a nil fake IP manager. internalDnatFw only checked for a firewall and GOOS, so the interceptor took the DNAT path and called GetFakeIP/AllocateFakeIP on the nil *fakeip.Manager. These methods lock m.mu first, which is a nil pointer dereference: the first DNS answer for such a route panicked and crashed the VPN service. The fake IP blocks (240.0.0.0/8 and its v6 pair) also never reached the TUN, since only the constructor registered them. Create the manager and its TUN routes from UpdateRoutes when the flag turns on, notify so the fake IP blocks get into the TUN without a client route change, and treat a nil manager as no internal DNAT. This is groundwork for removing the initial GetNetworkMap fetch, after which every startup goes through the flag-off-to-on transition. --- .../routemanager/dnsinterceptor/handler.go | 2 +- client/internal/routemanager/manager.go | 48 +++++++++++-------- .../routemanager/notifier/notifier_android.go | 1 + 3 files changed, 30 insertions(+), 21 deletions(-) diff --git a/client/internal/routemanager/dnsinterceptor/handler.go b/client/internal/routemanager/dnsinterceptor/handler.go index f92300bfd..d20f4b944 100644 --- a/client/internal/routemanager/dnsinterceptor/handler.go +++ b/client/internal/routemanager/dnsinterceptor/handler.go @@ -479,7 +479,7 @@ func (d *DnsInterceptor) removeDNATMappings(realPrefixes []netip.Prefix, logger // internalDnatFw checks if the firewall supports internal DNAT func (d *DnsInterceptor) internalDnatFw() (internalDNATer, bool) { - if d.firewall == nil || runtime.GOOS != "android" { + if d.firewall == nil || d.fakeIPManager == nil || runtime.GOOS != "android" { return nil, false } fw, ok := d.firewall.(internalDNATer) diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go index 2ab7e2a85..0cb74fd45 100644 --- a/client/internal/routemanager/manager.go +++ b/client/internal/routemanager/manager.go @@ -165,31 +165,36 @@ func (m *DefaultManager) setupAndroidRoutes(config ManagerConfig) { routesForComparison := slices.Clone(cr) if config.DNSFeatureFlag { - m.fakeIPManager = fakeip.NewManager() - - v4ID := uuid.NewString() - fakeIPRoute := &route.Route{ - ID: route.ID(v4ID), - Network: m.fakeIPManager.GetFakeIPBlock(), - NetID: route.NetID(v4ID), - Peer: m.pubKey, - NetworkType: route.IPv4Network, - } - v6ID := uuid.NewString() - fakeIPv6Route := &route.Route{ - ID: route.ID(v6ID), - Network: m.fakeIPManager.GetFakeIPv6Block(), - NetID: route.NetID(v6ID), - Peer: m.pubKey, - NetworkType: route.IPv6Network, - } - cr = append(cr, fakeIPRoute, fakeIPv6Route) - m.notifier.SetFakeIPRoutes([]*route.Route{fakeIPRoute, fakeIPv6Route}) + cr = append(cr, m.enableFakeIPRoutes()...) } m.notifier.SetInitialClientRoutes(cr, routesForComparison) } +func (m *DefaultManager) enableFakeIPRoutes() []*route.Route { + m.fakeIPManager = fakeip.NewManager() + + v4ID := uuid.NewString() + fakeIPRoute := &route.Route{ + ID: route.ID(v4ID), + Network: m.fakeIPManager.GetFakeIPBlock(), + NetID: route.NetID(v4ID), + Peer: m.pubKey, + NetworkType: route.IPv4Network, + } + v6ID := uuid.NewString() + fakeIPv6Route := &route.Route{ + ID: route.ID(v6ID), + Network: m.fakeIPManager.GetFakeIPv6Block(), + NetID: route.NetID(v6ID), + Peer: m.pubKey, + NetworkType: route.IPv6Network, + } + fakeRoutes := []*route.Route{fakeIPRoute, fakeIPv6Route} + m.notifier.SetFakeIPRoutes(fakeRoutes) + return fakeRoutes +} + func (m *DefaultManager) setupRefCounters(useNoop bool) { var once sync.Once var wgIface *net.Interface @@ -464,6 +469,9 @@ func (m *DefaultManager) UpdateRoutes( var merr *multierror.Error if !m.disableClientRoutes { + if runtime.GOOS == "android" && useNewDNSRoute && m.fakeIPManager == nil { + m.enableFakeIPRoutes() + } // Update route selector based on management server's isSelected status m.updateRouteSelectorFromManagement(clientRoutes) diff --git a/client/internal/routemanager/notifier/notifier_android.go b/client/internal/routemanager/notifier/notifier_android.go index 49300dbb2..31f2a7aef 100644 --- a/client/internal/routemanager/notifier/notifier_android.go +++ b/client/internal/routemanager/notifier/notifier_android.go @@ -41,6 +41,7 @@ func (n *Notifier) SetInitialClientRoutes(initialRoutes []*route.Route, routesFo // SetFakeIPRoutes stores the fake IP routes to be included in every TUN rebuild. func (n *Notifier) SetFakeIPRoutes(routes []*route.Route) { n.fakeIPRoutes = routes + n.notify() } func (n *Notifier) OnNewRoutes(idMap route.HAMap) { From 6f42636514c8d8483c23b267d9acf79a5e43c0c7 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 15:03:15 +0200 Subject: [PATCH 19/34] [client] Android - Serialize Android tunnel reconfiguration callbacks (#6990) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes The Android route notifier and the DNS search-domain notifier both delivered OnNetworkChanged from a fire-and-forget goroutine per update. Two updates in quick succession could reach the Java side reordered: the TUN rebuild handler applies them in arrival order and compares against the last applied parameters, so a stale route set delivered last won as the final TUN state. This is the same reordering hazard fixed for iOS in #6454. Wrap the Android network change listener into the shared tunnelnotifier FIFO introduced in #6870, the same way RunOniOS does, and deliver both notifiers synchronously into it. Enqueueing is non-blocking, a single delivery goroutine preserves order, and calls into Java never overlap. Also stop hasRouteDiff from sorting the notifier's shared route slices in place; compare sorted copies instead. ## 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/__ --- client/internal/connect.go | 5 ++++- client/internal/dns/notifier.go | 4 +--- .../routemanager/notifier/notifier_android.go | 19 ++++++------------- 3 files changed, 11 insertions(+), 17 deletions(-) diff --git a/client/internal/connect.go b/client/internal/connect.go index 87126b222..ceb39419e 100644 --- a/client/internal/connect.go +++ b/client/internal/connect.go @@ -113,11 +113,14 @@ func (c *ConnectClient) RunOnAndroid( stateFilePath string, cacheDir string, ) error { + notifier := tunnelnotifier.New(networkChangeListener, nil) + defer notifier.Close() + // in case of non Android os these variables will be nil mobileDependency := MobileDependency{ TunAdapter: tunAdapter, IFaceDiscover: iFaceDiscover, - NetworkChangeListener: networkChangeListener, + NetworkChangeListener: notifier, HostDNSAddresses: dnsAddresses, DnsReadyListener: dnsReadyListener, StateFilePath: stateFilePath, diff --git a/client/internal/dns/notifier.go b/client/internal/dns/notifier.go index 35cb6ff82..79d924a78 100644 --- a/client/internal/dns/notifier.go +++ b/client/internal/dns/notifier.go @@ -51,7 +51,5 @@ func (n *notifier) notify() { return } - go func(l listener.NetworkChangeListener) { - l.OnNetworkChanged("") - }(n.listener) + n.listener.OnNetworkChanged("") } diff --git a/client/internal/routemanager/notifier/notifier_android.go b/client/internal/routemanager/notifier/notifier_android.go index 31f2a7aef..60e1d0a0f 100644 --- a/client/internal/routemanager/notifier/notifier_android.go +++ b/client/internal/routemanager/notifier/notifier_android.go @@ -79,9 +79,7 @@ func (n *Notifier) notify() { routeStrings := n.routesToStrings(allRoutes) sort.Strings(routeStrings) - go func(l listener.NetworkChangeListener) { - l.OnNetworkChanged(strings.Join(routeStrings, ",")) - }(n.listener) + n.listener.OnNetworkChanged(strings.Join(routeStrings, ",")) } func filterStatic(routes []*route.Route) []*route.Route { @@ -103,16 +101,11 @@ func (n *Notifier) routesToStrings(routes []*route.Route) []string { } func (n *Notifier) hasRouteDiff(a []*route.Route, b []*route.Route) bool { - slices.SortFunc(a, func(x, y *route.Route) int { - return strings.Compare(x.NetString(), y.NetString()) - }) - slices.SortFunc(b, func(x, y *route.Route) int { - return strings.Compare(x.NetString(), y.NetString()) - }) - - return !slices.EqualFunc(a, b, func(x, y *route.Route) bool { - return x.NetString() == y.NetString() - }) + as := n.routesToStrings(a) + bs := n.routesToStrings(b) + sort.Strings(as) + sort.Strings(bs) + return !slices.Equal(as, bs) } func (n *Notifier) GetInitialRouteRanges() []string { From ee1389d736a442a098cb798346a92f3129c7f051 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Mon, 3 Aug 2026 23:23:10 +0900 Subject: [PATCH 20/34] [client] Keep the UI running when the notification service fails to start (#6959) --- client/ui/main.go | 7 ++- client/ui/notifier.go | 101 +++++++++++++++++++++++++++++++++++++++ client/ui/tray.go | 2 +- client/ui/tray_notify.go | 2 +- client/ui/tray_update.go | 4 +- 5 files changed, 108 insertions(+), 8 deletions(-) create mode 100644 client/ui/notifier.go diff --git a/client/ui/main.go b/client/ui/main.go index 782e76d1c..e2d172e5b 100644 --- a/client/ui/main.go +++ b/client/ui/main.go @@ -14,7 +14,6 @@ import ( "github.com/sirupsen/logrus" "github.com/wailsapp/wails/v3/pkg/application" "github.com/wailsapp/wails/v3/pkg/events" - "github.com/wailsapp/wails/v3/pkg/services/notifications" "github.com/netbirdio/netbird/client/ui/authsession" "github.com/netbirdio/netbird/client/ui/i18n" @@ -63,7 +62,7 @@ type registeredServices struct { profiles *services.Profiles update *services.Update daemonFeed *services.DaemonFeed - notifier *notifications.NotificationService + notifier *Notifier compat *services.Compat profileSwitcher *services.ProfileSwitcher bundle *i18n.Bundle @@ -102,7 +101,7 @@ func main() { updaterHolder := updater.NewHolder(app.Event) update := services.NewUpdate(conn, updaterHolder) daemonFeed := services.NewDaemonFeed(conn, app.Event, updaterHolder, debugLog) - notifier := notifications.New() + notifier := newNotifier() compat := services.NewCompat(conn) // macOS shows no toast until permission is requested. Run it after // ApplicationStarted so the notifier's Startup has initialised the @@ -210,7 +209,7 @@ func main() { // requestNotificationAuthorization prompts for macOS notification permission. // The request blocks until the user responds (up to 3 minutes), so callers run // it in a goroutine. No-op on Linux/Windows. -func requestNotificationAuthorization(notifier *notifications.NotificationService) { +func requestNotificationAuthorization(notifier *Notifier) { authorized, err := notifier.CheckNotificationAuthorization() if err != nil { logrus.Debugf("check notification authorization: %v", err) diff --git a/client/ui/notifier.go b/client/ui/notifier.go new file mode 100644 index 000000000..71ae3b0df --- /dev/null +++ b/client/ui/notifier.go @@ -0,0 +1,101 @@ +//go:build !android && !ios && !freebsd && !js + +package main + +import ( + "context" + "errors" + "sync/atomic" + + log "github.com/sirupsen/logrus" + "github.com/wailsapp/wails/v3/pkg/application" + "github.com/wailsapp/wails/v3/pkg/services/notifications" +) + +var errNotificationsUnavailable = errors.New("notifications unavailable") + +// Notifier wraps the Wails notification service so an unavailable backend +// disables notifications instead of aborting the app. Startup fails for +// environment reasons (a bare unbundled binary on macOS has no bundle +// identifier, a headless Linux session has no D-Bus session bus), and Wails +// treats a service startup error as fatal. After a failed startup every call +// is a no-op: on macOS, touching UNUserNotificationCenter without a bundle +// identifier raises an Objective-C exception that recover() cannot catch. +type Notifier struct { + inner *notifications.NotificationService + available atomic.Bool +} + +func newNotifier() *Notifier { + return &Notifier{inner: notifications.New()} +} + +// ServiceName implements the Wails service-name hook for startup logs. +func (n *Notifier) ServiceName() string { + return n.inner.ServiceName() +} + +// ServiceStartup starts the platform notifier, downgrading failure to a +// warning so the app keeps running without notifications. +func (n *Notifier) ServiceStartup(ctx context.Context, options application.ServiceOptions) error { + if err := n.inner.ServiceStartup(ctx, options); err != nil { + log.Warnf("notifications disabled: %v", err) + return nil + } + n.available.Store(true) + return nil +} + +func (n *Notifier) ServiceShutdown() error { + if !n.available.Load() { + return nil + } + return n.inner.ServiceShutdown() +} + +func (n *Notifier) CheckNotificationAuthorization() (bool, error) { + if !n.available.Load() { + return false, errNotificationsUnavailable + } + return n.inner.CheckNotificationAuthorization() +} + +func (n *Notifier) RequestNotificationAuthorization() (bool, error) { + if !n.available.Load() { + return false, errNotificationsUnavailable + } + return n.inner.RequestNotificationAuthorization() +} + +// SendNotification delivers a notification, silently dropping it when the +// backend never started (notifications are best-effort everywhere). +func (n *Notifier) SendNotification(options notifications.NotificationOptions) error { + if !n.available.Load() { + log.Debugf("notifications disabled, dropping %q", options.ID) + return nil + } + return n.inner.SendNotification(options) +} + +func (n *Notifier) SendNotificationWithActions(options notifications.NotificationOptions) error { + if !n.available.Load() { + log.Debugf("notifications disabled, dropping %q", options.ID) + return nil + } + return n.inner.SendNotificationWithActions(options) +} + +func (n *Notifier) RegisterNotificationCategory(category notifications.NotificationCategory) error { + if !n.available.Load() { + return nil + } + return n.inner.RegisterNotificationCategory(category) +} + +// OnNotificationResponse registers the response callback. Pure Go state, so +// it is safe (and simply inert) when the backend never started. +// +//wails:ignore +func (n *Notifier) OnNotificationResponse(callback func(result notifications.NotificationResult)) { + n.inner.OnNotificationResponse(callback) +} diff --git a/client/ui/tray.go b/client/ui/tray.go index 3050d159a..3093c693b 100644 --- a/client/ui/tray.go +++ b/client/ui/tray.go @@ -44,7 +44,7 @@ type TrayServices struct { Profiles *services.Profiles Networks *services.Networks DaemonFeed *services.DaemonFeed - Notifier *notifications.NotificationService + Notifier *Notifier Update *services.Update ProfileSwitcher *services.ProfileSwitcher WindowManager *services.WindowManager diff --git a/client/ui/tray_notify.go b/client/ui/tray_notify.go index d1117b57b..5b2629419 100644 --- a/client/ui/tray_notify.go +++ b/client/ui/tray_notify.go @@ -44,7 +44,7 @@ func safeSendNotification(send sendFn, what string, opts notifications.Notificat // notifyIfDaemonOutdated probes the daemon once and fires an OS toast when it // is reachable but too old for this UI. A probe error means the daemon isn't // reachable (not outdated), so it is left to the normal connection flow. -func notifyIfDaemonOutdated(compat *services.Compat, notifier *notifications.NotificationService, loc *Localizer) { +func notifyIfDaemonOutdated(compat *services.Compat, notifier *Notifier, loc *Localizer) { ready, err := compat.DaemonReady(context.Background()) if err != nil { log.Debugf("daemon compatibility probe: %v", err) diff --git a/client/ui/tray_update.go b/client/ui/tray_update.go index 1a377dfa3..27037eccb 100644 --- a/client/ui/tray_update.go +++ b/client/ui/tray_update.go @@ -21,7 +21,7 @@ type trayUpdater struct { app *application.App window *application.WebviewWindow update *services.Update - notifier *notifications.NotificationService + notifier *Notifier loc *Localizer onIconChange func() // onMenuChange drives a full tray relayout: the update row lives in the @@ -36,7 +36,7 @@ type trayUpdater struct { progressWindowOpen bool } -func newTrayUpdater(app *application.App, window *application.WebviewWindow, update *services.Update, notifier *notifications.NotificationService, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { +func newTrayUpdater(app *application.App, window *application.WebviewWindow, update *services.Update, notifier *Notifier, loc *Localizer, onIconChange func(), onMenuChange func()) *trayUpdater { u := &trayUpdater{ app: app, window: window, From d29bc23bb7fe62604a021438b2ea7acf7ae4a42b Mon Sep 17 00:00:00 2001 From: Riccardo Manfrin <3090891+riccardomanfrin@users.noreply.github.com> Date: Mon, 3 Aug 2026 16:28:10 +0200 Subject: [PATCH 21/34] [client] launch macOS GUI as the logged-in user after install/update (#6962) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes In unattended installs/updates there may be no logged-in user, so there's no context to start the GUI (nor anyone to see it). The bug is the GUI being started in the wrong user context / inheriting the wrong `$HOME` (OS mechanics aside). Today the GUI is started by a per-user LaunchAgent, i.e. on behalf of the user who logs in — no login ⇒ no GUI. The patch aligns to this: it launches the GUI on behalf of the logged-in console user if one exists, otherwise it delegates the launch to the per-user LaunchAgent at next login. Additionally it logs when default UI settings are applied. Note (small caveat): the LaunchAgent auto-starts the GUI at login only once it's been registered — which happens on the first GUI launch in the user's context. On an MDM/unattended fresh install done with no user logged in (where the user has never run the GUI before), they may need to start it manually once; it self-registers from then on. ## Issue ticket number and link No public issue — reported internally (community report on Slack: macOS advanced-view + onboarding reset on every update, esp. via MDM/Munki). Buggy line on main: https://github.com/netbirdio/netbird/blob/dd2bdc0de3aa14dd14dc90611c69c420bb3d3eb2/release_files/darwin_pkg/postinstall#L33 ## 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 macOS installer / GUI-launch behavior. No public API, CLI, or configuration change: the fix only changes the user context the desktop GUI is launched in after a pkg install/update. ### 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 macOS installation and updater UI launching to occur only when an active, valid GUI console session is detected. - Prevented UI launches during unattended/system, root, or login-window-related installs. - Ensured the app is launched in the correct console-user context, and skips cleanly when username/UID resolution fails. - **Improvements** - Added clearer informational logging when the UI preferences file is not found and default preferences are used. --- .../updater/installer/installer_run_darwin.go | 49 +++++++++---------- client/ui/preferences/store.go | 1 + release_files/darwin_pkg/postinstall | 18 ++++++- 3 files changed, 41 insertions(+), 27 deletions(-) diff --git a/client/internal/updater/installer/installer_run_darwin.go b/client/internal/updater/installer/installer_run_darwin.go index 248a404aa..5650bc769 100644 --- a/client/internal/updater/installer/installer_run_darwin.go +++ b/client/internal/updater/installer/installer_run_darwin.go @@ -98,47 +98,44 @@ func (u *Installer) startDaemon(daemonFolder string) error { func (u *Installer) startUIAsUser() error { log.Infof("starting netbird-ui: %s", uiBinary) - // Get the current console user - cmd := exec.Command("stat", "-f", "%Su", "/dev/console") - output, err := cmd.Output() + username, err := consoleUser() if err != nil { - return fmt.Errorf("failed to get console user: %w", err) + return err } - username := strings.TrimSpace(string(output)) - if username == "" || username == "root" { - return fmt.Errorf("no active user session found") - } - - log.Infof("starting UI for user: %s", username) - - // Get user's UID userInfo, err := user.Lookup(username) if err != nil { - return fmt.Errorf("failed to lookup user %s: %w", username, err) + return fmt.Errorf("lookup user %s: %w", username, err) } - // Start the UI process as the console user using launchctl - // This ensures the app runs in the user's context with proper GUI access - launchCmd := exec.Command("launchctl", "asuser", userInfo.Uid, "open", "-a", uiBinary) + log.Infof("starting UI for user: %s (uid %s)", username, userInfo.Uid) + + launchCmd := exec.Command("launchctl", "asuser", userInfo.Uid, "sudo", "-u", username, "-H", "open", "-a", uiBinary) log.Infof("launchCmd: %s", launchCmd.String()) - // Set the user's home directory for proper macOS app behavior - launchCmd.Env = append(os.Environ(), "HOME="+userInfo.HomeDir) - log.Infof("set HOME environment variable: %s", userInfo.HomeDir) - if err := launchCmd.Start(); err != nil { - return fmt.Errorf("failed to start UI process: %w", err) - } - - // Release the process so it can run independently - if err := launchCmd.Process.Release(); err != nil { - log.Warnf("failed to release UI process: %v", err) + if err := launchCmd.Run(); err != nil { + return fmt.Errorf("run UI launch: %w", err) } log.Infof("netbird-ui started successfully for user %s", username) return nil } +func consoleUser() (string, error) { + output, err := exec.Command("stat", "-f", "%Su", "/dev/console").Output() + if err != nil { + return "", fmt.Errorf("get console user: %w", err) + } + + username := strings.TrimSpace(string(output)) + switch username { + case "", "root", "loginwindow", "_mbsetupuser": + return "", fmt.Errorf("no active GUI user session, console user: %q", username) + } + + return username, nil +} + func (u *Installer) installPkgFile(ctx context.Context, path string) error { log.Infof("installing pkg file: %s", path) diff --git a/client/ui/preferences/store.go b/client/ui/preferences/store.go index df6fbbb16..49acb7917 100644 --- a/client/ui/preferences/store.go +++ b/client/ui/preferences/store.go @@ -246,6 +246,7 @@ func (s *Store) ExistedAtLoad() bool { func (s *Store) load() error { if _, err := os.Stat(s.path); err != nil { if errors.Is(err, os.ErrNotExist) { + log.Infof("no ui preferences file at %s; using defaults", s.path) return nil } return fmt.Errorf("stat preferences: %w", err) diff --git a/release_files/darwin_pkg/postinstall b/release_files/darwin_pkg/postinstall index 33fa4bfee..2c96a80cd 100755 --- a/release_files/darwin_pkg/postinstall +++ b/release_files/darwin_pkg/postinstall @@ -30,7 +30,23 @@ mkdir -p /usr/local/bin/ $AGENT service install || true $AGENT service start || true - open $APP + console_user=$(stat -f%Su /dev/console 2>/dev/null) + case "$console_user" in + ""|root|loginwindow|_mbsetupuser) + echo "No active GUI user session (console user: '${console_user:-none}'); skipping UI launch." + ;; + *) + uid=$(id -u "$console_user" 2>/dev/null) + if [ -z "$uid" ]; then + echo "Could not resolve uid for console user '$console_user'; skipping UI launch." + else + echo "Launching NetBird UI as console user $console_user (uid $uid)." + if ! launchctl asuser "$uid" sudo -u "$console_user" -H open "$APP"; then + echo "Failed to launch NetBird UI; if autostart is enabled it will start at next login." + fi + fi + ;; + esac echo "Finished Netbird installation successfully" exit 0 # all good From e90be36cd51afe7c039b673d4b421dc5c36a4700 Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Mon, 3 Aug 2026 23:28:51 +0900 Subject: [PATCH 22/34] [client] Don't ask for an SSO login when the login never reached management (#6983) --- client/server/login_outcome_test.go | 89 +++++++++++++++++++++++++++++ client/server/server.go | 37 ++++++++++-- 2 files changed, 122 insertions(+), 4 deletions(-) create mode 100644 client/server/login_outcome_test.go diff --git a/client/server/login_outcome_test.go b/client/server/login_outcome_test.go new file mode 100644 index 000000000..d3b1b7c52 --- /dev/null +++ b/client/server/login_outcome_test.go @@ -0,0 +1,89 @@ +package server + +import ( + "context" + "encoding/json" + "errors" + "os" + "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/proto" +) + +// A login that never reached Management is not a decision about the peer's +// credentials, so it must come back as a retryable error rather than an SSO +// prompt: the user cannot finish a browser login while Management is down, and +// the CLI's own backoff resolves the outage on its own once the daemon reports +// the failure. Reproduces `netbird down; netbird up` printing a device-code URL +// because Management happened to be restarting when the daemon dialed it. +func TestLogin_ManagementUnreachableIsReturnedInsteadOfDemandingSSO(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + unreachable := errors.New("create connection: dial context: context deadline exceeded") + attempts := 0 + s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) { + attempts++ + return internal.StatusLoginFailed, unreachable + } + + resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) + require.Error(t, err) + require.ErrorIs(t, err, unreachable, "the transport failure was replaced by something else") + require.Nil(t, resp, "a failed login must not answer with a login response") + require.Equal(t, 1, attempts) + require.Nil(t, s.oauthAuthFlow.flow, "the daemon started an SSO flow for a peer whose login was never decided") + + status, err := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, err) + require.Equal(t, internal.StatusLoginFailed, status, + "a peer that could not reach Management is not waiting on a login") +} + +// The counterpart: Management refusing the peer's credentials is a decision, and +// the SSO flow still has to start for it. The profile carries an unusable +// private key so the flow setup fails immediately instead of dialing, which is +// enough to show the branch was entered — the refusal itself is never what comes +// back out. +func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) { + s, _, _, username, cfgPath := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + breakProfilePrivateKey(t, cfgPath) + + refused := gstatus.Error(codes.PermissionDenied, "peer is not registered") + s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) { + return internal.StatusNeedsLogin, refused + } + + _, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) + require.Error(t, err) + require.NotErrorIs(t, err, refused, + "the refusal was handed back to the caller instead of starting the SSO flow") + + status, stateErr := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, stateErr) + require.Equal(t, internal.StatusLoginFailed, status, + "the SSO flow setup was never reached with the broken key") +} + +// breakProfilePrivateKey replaces the profile's private key with an unparseable +// one, which makes any attempt to build a Management client fail on the spot. +func breakProfilePrivateKey(t *testing.T, cfgPath string) { + t.Helper() + + raw, err := os.ReadFile(cfgPath) + require.NoError(t, err) + + var cfg map[string]any + require.NoError(t, json.Unmarshal(raw, &cfg)) + cfg["PrivateKey"] = "not-a-key" + + patched, err := json.Marshal(cfg) + require.NoError(t, err) + require.NoError(t, os.WriteFile(cfgPath, patched, 0o600)) +} diff --git a/client/server/server.go b/client/server/server.go index aaab5cc02..892c9c5de 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -135,6 +135,11 @@ type Server struct { updateManager *updater.Manager jwtCache *jwtCache + + // loginAttemptFn stands in for the Management login round trip. Tests set + // it to drive the login outcomes that need a server on the other end; + // production leaves it nil, and every login goes through loginAttempt. + loginAttemptFn func(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) } type oauthAuthFlow struct { @@ -370,7 +375,19 @@ func (s *Server) connectionGoroutineRunning() bool { } } -// loginAttempt attempts to login using the provided information. it returns a status in case something fails +// attemptLogin runs a login round trip against Management, or the stand-in a +// test installed in place of it. +func (s *Server) attemptLogin(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) { + if s.loginAttemptFn != nil { + return s.loginAttemptFn(ctx, setupKey, jwtToken) + } + return s.loginAttempt(ctx, setupKey, jwtToken) +} + +// loginAttempt attempts to login using the provided information. It returns +// StatusNeedsLogin when Management refused the peer's credentials and +// StatusLoginFailed for every other failure, so callers can tell an +// authentication decision apart from a login that never got made. func (s *Server) loginAttempt(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) { authClient, err := auth.NewAuth(ctx, s.config.PrivateKey, s.config.ManagementURL, s.config) if err != nil { @@ -623,11 +640,23 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro s.config = config s.mutex.Unlock() - if _, err := s.loginAttempt(ctx, "", ""); err == nil { + loginStatus, err := s.attemptLogin(ctx, "", "") + if err == nil { state.Set(internal.StatusIdle) return &proto.LoginResponse{}, nil } + // Only an authentication refusal means the peer has to (re-)authenticate. + // Any other failure leaves the login undecided: Management unreachable, a + // restart mid-request, an internal error. Those are returned for the caller + // to retry, because turning them into an SSO prompt asks the user to solve + // something that is not theirs to solve, and a browser login cannot succeed + // while Management is unreachable anyway. + if loginStatus != internal.StatusNeedsLogin { + state.Set(loginStatus) + return nil, err + } + if msg.SetupKey == "" { hint := "" if msg.Hint != nil { @@ -684,7 +713,7 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro // which returns NeedsLogin and parks on the browser leg. state.Set(internal.StatusConnecting) - if loginStatus, err := s.loginAttempt(ctx, msg.SetupKey, ""); err != nil { + if loginStatus, err := s.attemptLogin(ctx, msg.SetupKey, ""); err != nil { state.Set(loginStatus) return nil, err } @@ -839,7 +868,7 @@ func (s *Server) WaitSSOLogin(callerCtx context.Context, msg *proto.WaitSSOLogin s.oauthAuthFlow.expiresAt = time.Now() s.mutex.Unlock() - if loginStatus, err := s.loginAttempt(ctx, "", tokenInfo.GetTokenToUse()); err != nil { + if loginStatus, err := s.attemptLogin(ctx, "", tokenInfo.GetTokenToUse()); err != nil { state.Set(loginStatus) return nil, err } From b82a42c855bf82d02405ea461b2f6455c648eb26 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 17:58:29 +0200 Subject: [PATCH 23/34] [misc] Add android and ios tags to PR title check (#7037) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Extend CI tag list with Android and iOS ## 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) - [ ] 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). ## 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 * **Chores** * Updated pull request title validation to accept `android` and `ios` tags. --- .github/workflows/pr-title-check.yml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.github/workflows/pr-title-check.yml b/.github/workflows/pr-title-check.yml index 67d65356c..24d81b50f 100644 --- a/.github/workflows/pr-title-check.yml +++ b/.github/workflows/pr-title-check.yml @@ -16,6 +16,8 @@ jobs: const allowedTags = [ 'management', 'client', + 'android', + 'ios', 'signal', 'proxy', 'relay', From 075b319fb34d83d7089da05ca42d0dbfb6f705ab Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 18:23:02 +0200 Subject: [PATCH 24/34] [client, android] Pull fresh TUN settings on Android rebuild (#6991) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Pull fresh TUN settings on Android rebuild instead of push The Android TUN rebuild consumed state pushed through notifications and a Java-side snapshot, and both sources were unreliable. The DNS search-domain notifier fired OnNetworkChanged with an empty string, which the rebuild handler treated as the new route list, so any search domain change rebuilt the TUN with zero routes and cut all tunnel traffic. The rebuild also reused the search domains cached at the last establish, so search domain updates never reached the TUN at runtime. Make the notification a pure trigger and let the Java side pull a fresh snapshot instead. Expose GetTunSettings on the Android SDK client: it returns the current TUN route ranges, derived on demand by the route manager from the client routes, the exit-node selection and the fake IP blocks, together with the DNS search domains. The route notifier keeps only its last-announced baseline to suppress triggers for unchanged syncs; the TUN route state is owned by the route manager. SearchDomains now locks the DNS server mutex since the pull arrives from a Java thread. Requires the matching android-client change that switches recreateTUN to the pull API. ## 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/__ ## Summary by CodeRabbit - **New Features** - Added access to current TUN route ranges and DNS search domains. - TUN settings are returned in a mobile-friendly format for easier integration. - **Improvements** - Route changes are detected and synchronized more reliably. - Current routing information now reflects active routes, including supported fake-IP ranges. - Simplified network initialization for more consistent startup behavior. - **API Changes** - Removed the obsolete network-map retrieval method from the management client interface. --- client/android/client.go | 24 +++++ client/internal/dns/server.go | 10 ++- client/internal/engine.go | 51 +---------- client/internal/engine_tunsettings.go | 20 +++++ client/internal/routemanager/manager.go | 88 +++++++------------ client/internal/routemanager/mock.go | 4 +- .../routemanager/notifier/notifier_android.go | 78 ++++++---------- .../routemanager/notifier/notifier_ios.go | 6 +- .../routemanager/notifier/notifier_other.go | 10 +-- shared/management/client/client.go | 1 - shared/management/client/grpc.go | 43 --------- shared/management/client/mock.go | 5 -- 12 files changed, 116 insertions(+), 224 deletions(-) create mode 100644 client/internal/engine_tunsettings.go diff --git a/client/android/client.go b/client/android/client.go index 501d7f77c..59dcb1021 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -57,6 +57,12 @@ type DnsReadyListener interface { dns.ReadyListener } +// TunSettings is a snapshot of the settings the TUN device is rebuilt with +type TunSettings struct { + Routes string + SearchDomains string +} + func init() { formatter.SetLogcatFormatter(log.StandardLogger()) } @@ -240,6 +246,24 @@ func (c *Client) RenewTun(fd int) error { return e.RenewTun(fd) } +func (c *Client) GetTunSettings() (*TunSettings, error) { + cc := c.getConnectClient() + if cc == nil { + return nil, fmt.Errorf("engine not running") + } + + e := cc.Engine() + if e == nil { + return nil, fmt.Errorf("engine not initialized") + } + + routes, searchDomains := e.TunSettings() + return &TunSettings{ + Routes: strings.Join(routes, ";"), + SearchDomains: strings.Join(searchDomains, ";"), + }, nil +} + // DebugBundle generates a debug bundle, uploads it, and returns the upload key. // It works both with and without a running engine. func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (string, error) { diff --git a/client/internal/dns/server.go b/client/internal/dns/server.go index f79454457..3af912792 100644 --- a/client/internal/dns/server.go +++ b/client/internal/dns/server.go @@ -252,7 +252,7 @@ func NewDefaultServerPermanentUpstream( ds.hostsDNSHolder.set(hostsDnsList) ds.permanent = true ds.currentConfig = dnsConfigToHostDNSConfig(config, ds.service.RuntimeIP(), ds.service.RuntimePort()) - ds.searchDomainNotifier = newNotifier(ds.SearchDomains()) + ds.searchDomainNotifier = newNotifier(ds.searchDomains()) ds.searchDomainNotifier.setListener(listener) setServerDns(ds) return ds @@ -602,6 +602,12 @@ func (s *DefaultServer) UpdateDNSServer(serial uint64, update nbdns.Config) erro } func (s *DefaultServer) SearchDomains() []string { + s.mux.Lock() + defer s.mux.Unlock() + return s.searchDomains() +} + +func (s *DefaultServer) searchDomains() []string { var searchDomains []string for _, dConf := range s.currentConfig.Domains { @@ -686,7 +692,7 @@ func (s *DefaultServer) applyConfiguration(update nbdns.Config) error { }() if s.searchDomainNotifier != nil { - s.searchDomainNotifier.onNewSearchDomains(s.SearchDomains()) + s.searchDomainNotifier.onNewSearchDomains(s.searchDomains()) } s.updateNSGroupStates(update.NameServerGroups) diff --git a/client/internal/engine.go b/client/internal/engine.go index 617892e43..f4f47992f 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -572,12 +572,7 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) } e.stateManager.Start() - initialRoutes, dnsConfig, dnsFeatureFlag, err := e.readInitialSettings() - if err != nil { - return fmt.Errorf("read initial settings: %w", err) - } - - dnsServer, err := e.newDnsServer(dnsConfig) + dnsServer, err := e.newDnsServer() if err != nil { return fmt.Errorf("create dns server: %w", err) } @@ -595,10 +590,8 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL) WGInterface: e.wgInterface, StatusRecorder: e.statusRecorder, RelayManager: e.relayManager, - InitialRoutes: initialRoutes, StateManager: e.stateManager, DNSServer: dnsServer, - DNSFeatureFlag: dnsFeatureFlag, PeerStore: e.peerStore, DisableClientRoutes: e.config.DisableClientRoutes, DisableServerRoutes: e.config.DisableServerRoutes, @@ -2102,42 +2095,6 @@ func (e *Engine) close() { } } -func (e *Engine) readInitialSettings() ([]*route.Route, *nbdns.Config, bool, error) { - if runtime.GOOS != "android" { - // nolint:nilnil - return nil, nil, false, nil - } - - info := system.GetInfo(e.ctx) - info.SetFlags( - e.config.RosenpassEnabled, - e.config.RosenpassPermissive, - &e.config.ServerSSHAllowed, - e.config.DisableClientRoutes, - e.config.DisableServerRoutes, - e.config.DisableDNS, - e.config.DisableFirewall, - e.config.BlockLANAccess, - e.config.BlockInbound, - e.config.DisableIPv6, - e.config.SyncMessageVersion, - e.config.EnableSSHRoot, - e.config.EnableSSHSFTP, - e.config.EnableSSHLocalPortForwarding, - e.config.EnableSSHRemotePortForwarding, - e.config.DisableSSHAuth, - ) - - netMap, err := e.mgmClient.GetNetworkMap(info) - if err != nil { - return nil, nil, false, err - } - routes := toRoutes(netMap.GetRoutes()) - dnsCfg := toDNSConfig(netMap.GetDNSConfig(), e.wgInterface.Address()) - dnsFeatureFlag := toDNSFeatureFlag(netMap) - return routes, &dnsCfg, dnsFeatureFlag, nil -} - func (e *Engine) newWgIface() (*iface.WGIface, error) { transportNet, err := e.newStdNet() if err != nil { @@ -2172,7 +2129,7 @@ func (e *Engine) newWgIface() (*iface.WGIface, error) { func (e *Engine) wgInterfaceCreate() (err error) { switch runtime.GOOS { case "android": - err = e.wgInterface.CreateOnAndroid(e.routeManager.InitialRouteRange(), e.dnsServer.DnsIP().String(), e.dnsServer.SearchDomains()) + err = e.wgInterface.CreateOnAndroid(e.routeManager.CurrentRouteRange(), e.dnsServer.DnsIP().String(), e.dnsServer.SearchDomains()) case "ios": e.mobileDep.NetworkChangeListener.SetInterfaceIP(e.config.WgAddr.String()) if e.config.WgAddr.HasIPv6() { @@ -2185,7 +2142,7 @@ func (e *Engine) wgInterfaceCreate() (err error) { return err } -func (e *Engine) newDnsServer(dnsConfig *nbdns.Config) (dns.Server, error) { +func (e *Engine) newDnsServer() (dns.Server, error) { // due to tests where we are using a mocked version of the DNS server if e.dnsServer != nil { return e.dnsServer, nil @@ -2197,7 +2154,7 @@ func (e *Engine) newDnsServer(dnsConfig *nbdns.Config) (dns.Server, error) { e.ctx, e.wgInterface, e.mobileDep.HostDNSAddresses, - *dnsConfig, + nbdns.Config{}, e.mobileDep.NetworkChangeListener, e.statusRecorder, e.config.DisableDNS, diff --git a/client/internal/engine_tunsettings.go b/client/internal/engine_tunsettings.go new file mode 100644 index 000000000..34a59671a --- /dev/null +++ b/client/internal/engine_tunsettings.go @@ -0,0 +1,20 @@ +package internal + +func (e *Engine) TunSettings() ([]string, []string) { + e.syncMsgMux.Lock() + routeManager := e.routeManager + dnsServer := e.dnsServer + e.syncMsgMux.Unlock() + + var routes []string + if routeManager != nil { + routes = routeManager.CurrentRouteRange() + } + + var searchDomains []string + if dnsServer != nil { + searchDomains = dnsServer.SearchDomains() + } + + return routes, searchDomains +} diff --git a/client/internal/routemanager/manager.go b/client/internal/routemanager/manager.go index 0cb74fd45..0ccfa83ac 100644 --- a/client/internal/routemanager/manager.go +++ b/client/internal/routemanager/manager.go @@ -8,14 +8,13 @@ import ( "net/netip" "net/url" "runtime" - "slices" + "sort" "strings" "sync" "sync/atomic" "syscall" "time" - "github.com/google/uuid" "github.com/hashicorp/go-multierror" log "github.com/sirupsen/logrus" "golang.org/x/exp/maps" @@ -62,7 +61,7 @@ type Manager interface { GetActiveClientRoutes() route.HAMap GetClientRoutesWithNetID() map[route.NetID][]*route.Route SetRouteChangeListener(listener listener.NetworkChangeListener) - InitialRouteRange() []string + CurrentRouteRange() []string SetFirewall(firewall.Manager) error SetDNSForwarderPort(port uint16) ReconcilePeerAllowedIPs(peerKey string) error @@ -76,10 +75,8 @@ type ManagerConfig struct { WGInterface iface.WGIface StatusRecorder *peer.Status RelayManager *relayClient.Manager - InitialRoutes []*route.Route StateManager *statemanager.Manager DNSServer dns.Server - DNSFeatureFlag bool PeerStore *peerstore.Store DisableClientRoutes bool DisableServerRoutes bool @@ -149,50 +146,12 @@ func NewManager(config ManagerConfig) *DefaultManager { useNoop := netstack.IsEnabled() || config.DisableClientRoutes dm.setupRefCounters(useNoop) - // don't proceed with client routes if it is disabled - if config.DisableClientRoutes { - return dm - } - - if runtime.GOOS == "android" { - dm.setupAndroidRoutes(config) - } return dm } -func (m *DefaultManager) setupAndroidRoutes(config ManagerConfig) { - cr := m.initialClientRoutes(config.InitialRoutes) - routesForComparison := slices.Clone(cr) - - if config.DNSFeatureFlag { - cr = append(cr, m.enableFakeIPRoutes()...) - } - - m.notifier.SetInitialClientRoutes(cr, routesForComparison) -} - -func (m *DefaultManager) enableFakeIPRoutes() []*route.Route { +func (m *DefaultManager) enableFakeIPRoutes() { m.fakeIPManager = fakeip.NewManager() - - v4ID := uuid.NewString() - fakeIPRoute := &route.Route{ - ID: route.ID(v4ID), - Network: m.fakeIPManager.GetFakeIPBlock(), - NetID: route.NetID(v4ID), - Peer: m.pubKey, - NetworkType: route.IPv4Network, - } - v6ID := uuid.NewString() - fakeIPv6Route := &route.Route{ - ID: route.ID(v6ID), - Network: m.fakeIPManager.GetFakeIPv6Block(), - NetID: route.NetID(v6ID), - Peer: m.pubKey, - NetworkType: route.IPv6Network, - } - fakeRoutes := []*route.Route{fakeIPRoute, fakeIPv6Route} - m.notifier.SetFakeIPRoutes(fakeRoutes) - return fakeRoutes + m.notifier.NotifyRouteChange() } func (m *DefaultManager) setupRefCounters(useNoop bool) { @@ -508,9 +467,32 @@ func (m *DefaultManager) SetRouteChangeListener(listener listener.NetworkChangeL m.notifier.SetListener(listener) } -// InitialRouteRange return the list of initial routes. It used by mobile systems -func (m *DefaultManager) InitialRouteRange() []string { - return m.notifier.GetInitialRouteRanges() +// CurrentRouteRange returns the current TUN route list. It is used by mobile systems +func (m *DefaultManager) CurrentRouteRange() []string { + m.mux.Lock() + defer m.mux.Unlock() + + if m.disableClientRoutes { + return nil + } + + filtered := m.routeSelector.FilterSelectedExitNodes(m.clientRoutes) + var nets []string + for _, routes := range filtered { + for _, r := range routes { + if r.IsDynamic() { + continue + } + nets = append(nets, r.NetString()) + } + } + + if m.fakeIPManager != nil { + nets = append(nets, m.fakeIPManager.GetFakeIPBlock().String(), m.fakeIPManager.GetFakeIPv6Block().String()) + } + + sort.Strings(nets) + return nets } // GetRouteSelector returns the route selector @@ -708,16 +690,6 @@ func (m *DefaultManager) ClassifyRoutes(newRoutes []*route.Route) (map[route.ID] return newServerRoutesMap, newClientRoutesIDMap } -func (m *DefaultManager) initialClientRoutes(initialRoutes []*route.Route) []*route.Route { - _, crMap := m.ClassifyRoutes(initialRoutes) - rs := make([]*route.Route, 0, len(crMap)) - for _, routes := range crMap { - rs = append(rs, routes...) - } - - return rs -} - func isRouteSupported(route *route.Route) bool { if netstack.IsEnabled() || !nbnet.CustomRoutingDisabled() || route.IsDynamic() { return true diff --git a/client/internal/routemanager/mock.go b/client/internal/routemanager/mock.go index cf761091d..2a8398b95 100644 --- a/client/internal/routemanager/mock.go +++ b/client/internal/routemanager/mock.go @@ -30,8 +30,8 @@ func (m *MockManager) Init() error { return nil } -// InitialRouteRange mock implementation of InitialRouteRange from Manager interface -func (m *MockManager) InitialRouteRange() []string { +// CurrentRouteRange mock implementation of CurrentRouteRange from Manager interface +func (m *MockManager) CurrentRouteRange() []string { return nil } diff --git a/client/internal/routemanager/notifier/notifier_android.go b/client/internal/routemanager/notifier/notifier_android.go index 60e1d0a0f..5fa329310 100644 --- a/client/internal/routemanager/notifier/notifier_android.go +++ b/client/internal/routemanager/notifier/notifier_android.go @@ -6,7 +6,6 @@ import ( "net/netip" "slices" "sort" - "strings" "sync" "github.com/netbirdio/netbird/client/internal/listener" @@ -14,12 +13,15 @@ import ( ) type Notifier struct { - initialRoutes []*route.Route - currentRoutes []*route.Route - fakeIPRoutes []*route.Route + mu sync.Mutex - listener listener.NetworkChangeListener - listenerMux sync.Mutex + // currentRoutes is the last announced route set. It exists only to + // suppress noise: without it every network map sync would trigger the + // Java side, even when the routes did not change. The actual TUN route + // state is owned by the route manager and pulled from there. + currentRoutes []*route.Route + + listener listener.NetworkChangeListener } func NewNotifier() *Notifier { @@ -27,21 +29,15 @@ func NewNotifier() *Notifier { } func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { - n.listenerMux.Lock() - defer n.listenerMux.Unlock() + n.mu.Lock() + defer n.mu.Unlock() n.listener = listener } -// SetInitialClientRoutes stores the initial route sets for TUN configuration. -func (n *Notifier) SetInitialClientRoutes(initialRoutes []*route.Route, routesForComparison []*route.Route) { - n.initialRoutes = filterStatic(initialRoutes) - n.currentRoutes = filterStatic(routesForComparison) -} - -// SetFakeIPRoutes stores the fake IP routes to be included in every TUN rebuild. -func (n *Notifier) SetFakeIPRoutes(routes []*route.Route) { - n.fakeIPRoutes = routes - n.notify() +func (n *Notifier) NotifyRouteChange() { + n.mu.Lock() + defer n.mu.Unlock() + n.notifyLocked() } func (n *Notifier) OnNewRoutes(idMap route.HAMap) { @@ -55,44 +51,32 @@ func (n *Notifier) OnNewRoutes(idMap route.HAMap) { } } - if !n.hasRouteDiff(n.currentRoutes, newRoutes) { + n.mu.Lock() + defer n.mu.Unlock() + if !hasRouteDiff(n.currentRoutes, newRoutes) { return } n.currentRoutes = newRoutes - n.notify() + n.notifyLocked() } func (n *Notifier) OnNewPrefixes([]netip.Prefix) { // Not used on Android } -func (n *Notifier) notify() { - n.listenerMux.Lock() - defer n.listenerMux.Unlock() +func (n *Notifier) notifyLocked() { if n.listener == nil { return } - - allRoutes := slices.Clone(n.currentRoutes) - allRoutes = append(allRoutes, n.fakeIPRoutes...) - - routeStrings := n.routesToStrings(allRoutes) - sort.Strings(routeStrings) - n.listener.OnNetworkChanged(strings.Join(routeStrings, ",")) + n.listener.OnNetworkChanged("") } -func filterStatic(routes []*route.Route) []*route.Route { - out := make([]*route.Route, 0, len(routes)) - for _, r := range routes { - if !r.IsDynamic() { - out = append(out, r) - } - } - return out +func (n *Notifier) Close() { + // unused } -func (n *Notifier) routesToStrings(routes []*route.Route) []string { +func routesToStrings(routes []*route.Route) []string { nets := make([]string, 0, len(routes)) for _, r := range routes { nets = append(nets, r.NetString()) @@ -100,20 +84,10 @@ func (n *Notifier) routesToStrings(routes []*route.Route) []string { return nets } -func (n *Notifier) hasRouteDiff(a []*route.Route, b []*route.Route) bool { - as := n.routesToStrings(a) - bs := n.routesToStrings(b) +func hasRouteDiff(a []*route.Route, b []*route.Route) bool { + as := routesToStrings(a) + bs := routesToStrings(b) sort.Strings(as) sort.Strings(bs) return !slices.Equal(as, bs) } - -func (n *Notifier) GetInitialRouteRanges() []string { - initialStrings := n.routesToStrings(n.initialRoutes) - sort.Strings(initialStrings) - return initialStrings -} - -func (n *Notifier) Close() { - // unused -} diff --git a/client/internal/routemanager/notifier/notifier_ios.go b/client/internal/routemanager/notifier/notifier_ios.go index c91a76551..d663dd471 100644 --- a/client/internal/routemanager/notifier/notifier_ios.go +++ b/client/internal/routemanager/notifier/notifier_ios.go @@ -29,11 +29,7 @@ func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { n.listener = listener } -func (n *Notifier) SetInitialClientRoutes([]*route.Route, []*route.Route) { - // iOS doesn't care about initial routes -} - -func (n *Notifier) SetFakeIPRoutes([]*route.Route) { +func (n *Notifier) NotifyRouteChange() { // Not used on iOS } diff --git a/client/internal/routemanager/notifier/notifier_other.go b/client/internal/routemanager/notifier/notifier_other.go index 71b1096c2..fe48e07b3 100644 --- a/client/internal/routemanager/notifier/notifier_other.go +++ b/client/internal/routemanager/notifier/notifier_other.go @@ -19,11 +19,7 @@ func (n *Notifier) SetListener(listener listener.NetworkChangeListener) { // Not used on non-mobile platforms } -func (n *Notifier) SetInitialClientRoutes([]*route.Route, []*route.Route) { - // Not used on non-mobile platforms -} - -func (n *Notifier) SetFakeIPRoutes([]*route.Route) { +func (n *Notifier) NotifyRouteChange() { // Not used on non-mobile platforms } @@ -35,10 +31,6 @@ func (n *Notifier) OnNewPrefixes(prefixes []netip.Prefix) { // Not used on non-mobile platforms } -func (n *Notifier) GetInitialRouteRanges() []string { - return []string{} -} - func (n *Notifier) Close() { // unused } diff --git a/shared/management/client/client.go b/shared/management/client/client.go index 8205e3a4f..c48e1ed3e 100644 --- a/shared/management/client/client.go +++ b/shared/management/client/client.go @@ -22,7 +22,6 @@ type Client interface { ExtendAuthSession(sysInfo *system.Info, jwtToken string) (*proto.ExtendAuthSessionResponse, error) GetDeviceAuthorizationFlow() (*proto.DeviceAuthorizationFlow, error) GetPKCEAuthorizationFlow() (*proto.PKCEAuthorizationFlow, error) - GetNetworkMap(sysInfo *system.Info) (*proto.NetworkMap, error) GetServerURL() string // IsHealthy returns the current connection status without blocking. // Used by the engine to monitor connectivity in the background. diff --git a/shared/management/client/grpc.go b/shared/management/client/grpc.go index bd2d0da1f..81f25900a 100644 --- a/shared/management/client/grpc.go +++ b/shared/management/client/grpc.go @@ -436,49 +436,6 @@ func (c *GrpcClient) handleSyncStream(ctx context.Context, serverPubKey wgtypes. return nil } -// GetNetworkMap return with the network map -func (c *GrpcClient) GetNetworkMap(sysInfo *system.Info) (*proto.NetworkMap, error) { - serverPubKey, err := c.getServerPublicKey() - if err != nil { - log.Debugf("failed getting Management Service public key: %s", err) - return nil, err - } - - ctx, cancelStream := context.WithCancel(c.ctx) - defer cancelStream() - stream, err := c.connectToSyncStream(ctx, *serverPubKey, sysInfo) - if err != nil { - log.Debugf("failed to open Management Service stream: %s", err) - return nil, err - } - defer func() { - _ = stream.CloseSend() - }() - - update, err := stream.Recv() - if err == io.EOF { - log.Debugf("Management stream has been closed by server: %s", err) - return nil, err - } - if err != nil { - log.Debugf("disconnected from Management Service sync stream: %v", err) - return nil, err - } - - decryptedResp := &proto.SyncResponse{} - err = encryption.DecryptMessage(*serverPubKey, c.key, update.Body, decryptedResp) - if err != nil { - log.Errorf("failed decrypting update message from Management Service: %s", err) - return nil, err - } - - if decryptedResp.GetNetworkMap() == nil { - return nil, fmt.Errorf("invalid msg, required network map") - } - - return decryptedResp.GetNetworkMap(), nil -} - func (c *GrpcClient) connectToSyncStream(ctx context.Context, serverPubKey wgtypes.Key, sysInfo *system.Info) (proto.ManagementService_SyncClient, error) { req := &proto.SyncRequest{Meta: infoToMetaData(sysInfo)} diff --git a/shared/management/client/mock.go b/shared/management/client/mock.go index ba156a225..e57e314da 100644 --- a/shared/management/client/mock.go +++ b/shared/management/client/mock.go @@ -94,11 +94,6 @@ func (m *MockClient) HealthCheck() error { return m.HealthCheckFunc() } -// GetNetworkMap mock implementation of GetNetworkMap from Client interface. -func (m *MockClient) GetNetworkMap(_ *system.Info) (*proto.NetworkMap, error) { - return nil, nil -} - // GetServerURL mock implementation of GetServerURL from mgm.Client interface func (m *MockClient) GetServerURL() string { if m.GetServerURLFunc == nil { From 7546e7751cecffcf98f1b55e1feed88bd643da17 Mon Sep 17 00:00:00 2001 From: camiloariza <14282973+camiloariza@users.noreply.github.com> Date: Mon, 3 Aug 2026 13:42:15 -0300 Subject: [PATCH 25/34] [client, android] Reuse the persisted configuration when enrolling (#7022) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes NewAuth builds a fresh in-memory configuration on every call, which means a new WireGuard key each time. The peer registers under that key and the key is written out, so any peer registered by an earlier call is orphaned on the server — a client that enrols twice leaves two entries and owns neither. It also breaks the enrol-then-run sequence. `RunWithoutLogin` reloads the configuration from disk through `UpdateOrCreateConfig`, so the identity that registered is not necessarily the identity that runs, and the management stream rejects it: ``` failed to login to Management Service: rpc error: code = PermissionDenied desc = no peer auth method provided, please use a setup key or interactive SSO login ``` followed by a panic in `ConnectClient.run`. ### How it was found Embedding the Android client in an application that enrols with a setup key and then runs. Eight orphaned peers accumulated on a self-hosted management server before the cause was clear, because every restart registered a new one. ### The change `NewAuth` passes `ConfigPath` and uses `UpdateOrCreateConfig`, so an existing configuration is reused and one is only created when absent. A caller wanting a fresh identity can delete the file — which is what "forget this account" already does. ### Test `TestNewAuth_ReusesPersistedIdentity` fails on the current code: ``` --- FAIL: TestNewAuth_ReusesPersistedIdentity (0.00s) login_test.go:33: private key changed between calls: a second enrolment would orphan the peer registered by the first ``` and passes with the fix. `TestNewAuth_CreatesConfigWhenAbsent` covers the first-enrolment path being unchanged. Both run in `client/android` on Linux. Per CONTRIBUTING, opening directly as a bug fix rather than raising an issue first. ## Issue ticket number and link [NET-1465](https://linear.app/netbird/issue/NET-1465/agent-network-rest-api-settings-defaults-bootstrap-via-put-provider) ## 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) - [ ] 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). ## Documentation Select exactly one: - [ ] I added/updated documentation for this change - [x] Documentation is **not needed** for this change (explain why): the API reference is generated from the OpenAPI spec, which this PR updates in-repo. ### 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/__ Co-authored-by: Zoltan Papp --- client/android/login.go | 10 ++++++- client/android/login_test.go | 51 ++++++++++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 1 deletion(-) create mode 100644 client/android/login_test.go diff --git a/client/android/login.go b/client/android/login.go index a9422cdbf..32ce28739 100644 --- a/client/android/login.go +++ b/client/android/login.go @@ -36,12 +36,20 @@ type Auth struct { } // NewAuth instantiate Auth struct and validate the management URL +// +// The configuration at cfgPath is reused when one is already there, and only created when it is +// not. Building a fresh in-memory config unconditionally gives the client a new WireGuard key on +// every call: the peer registers under that key, the key is written out, and any peer registered by +// an earlier call is orphaned on the server. It also breaks a client that enrols and then runs from +// the persisted config, because the identity it registered is not the one it runs with — the +// management stream rejects it with "no peer auth method provided". func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { inputCfg := profilemanager.ConfigInput{ + ConfigPath: cfgPath, ManagementURL: mgmURL, } - cfg, err := profilemanager.CreateInMemoryConfig(inputCfg) + cfg, err := profilemanager.UpdateOrCreateConfig(inputCfg) if err != nil { return nil, err } diff --git a/client/android/login_test.go b/client/android/login_test.go new file mode 100644 index 000000000..b04790f6b --- /dev/null +++ b/client/android/login_test.go @@ -0,0 +1,51 @@ +package android + +import ( + "path/filepath" + "testing" +) + +// NewAuth must reuse the configuration already at cfgPath rather than building a fresh one. +// +// Creating a new in-memory config on every call gives the client a new WireGuard private key each +// time. The peer registers under that key and the key is written out, so a peer registered by an +// earlier call is orphaned on the server — a client that enrols twice leaves two entries and owns +// neither. It also breaks enrol-then-run: RunWithoutLogin reloads the configuration from disk, so +// the identity that registered is not the identity that runs, and the management stream rejects it +// with "no peer auth method provided, please use a setup key or interactive SSO login". +func TestNewAuth_ReusesPersistedIdentity(t *testing.T) { + cfgPath := filepath.Join(t.TempDir(), "config.json") + + first, err := NewAuth(cfgPath, "https://api.example.com:443") + if err != nil { + t.Fatalf("first NewAuth: %v", err) + } + if first.config.PrivateKey == "" { + t.Fatal("first NewAuth produced no private key") + } + + second, err := NewAuth(cfgPath, "https://api.example.com:443") + if err != nil { + t.Fatalf("second NewAuth: %v", err) + } + + if second.config.PrivateKey != first.config.PrivateKey { + t.Errorf("private key changed between calls: a second enrolment would orphan the peer registered by the first") + } +} + +// A missing configuration is still created, so a first enrolment works unchanged. +func TestNewAuth_CreatesConfigWhenAbsent(t *testing.T) { + cfgPath := filepath.Join(t.TempDir(), "config.json") + + auth, err := NewAuth(cfgPath, "https://api.example.com:443") + if err != nil { + t.Fatalf("NewAuth: %v", err) + } + if auth.config == nil || auth.config.PrivateKey == "" { + t.Fatal("NewAuth did not create a usable configuration") + } + if auth.cfgPath != cfgPath { + t.Errorf("cfgPath = %q, want %q", auth.cfgPath, cfgPath) + } +} From 1bedb4e59d2632dc301a574a61a3d67889c47041 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 20:50:57 +0000 Subject: [PATCH 26/34] [client, android] Reuse the profile's account for Android SSO logins (#6988) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Android binding never recorded which account a profile belongs to, so every interactive login and every session extend went to the IdP with no login_hint. With nothing to go on the IdP picks an account itself, which on a session extend means re-authenticating an account the profile is already signed in with. Store the email the PKCE flow already parses out of the ID token, and pass it back as the hint on later flows. An empty hint stays meaningful: a fresh profile, or one that was logged out, deliberately leaves the choice to the IdP, which is how a profile changes accounts. Logout clears the stored email for that reason — while it is on disk it would steer the next login straight back into the account just logged out of. The email is keyed off the profile's config path rather than the active profile: Auth.login runs in a goroutine, so the active profile can change under a flow already in flight. It lands in .account.json, not the .state.json desktop uses for the same data — there the email and the engine's state manager sit in different directories, but on Android both resolve under files/, and the state manager rewrites the whole file from its own keys. ## 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/__ ## Summary by CodeRabbit * **New Features** * Added account email and active-status details to Android profile information. * Improved SSO sign-in and session renewal by restoring the previously used account as a login hint. * Added Android-specific profile email persistence with automatic cleanup on logout. * **Bug Fixes** * Profile email persistence failures now generate warnings without blocking login or logout. * Improved handling of missing or unreadable account data and repeated logout cleanup. * **Tests** * Added coverage for account-file naming, email persistence, and logout behavior. --- client/android/client.go | 21 +++- client/android/login.go | 40 ++++++- client/android/profile_manager.go | 34 ++++-- client/android/profile_state.go | 108 ++++++++++++++++++ client/android/profile_state_test.go | 161 +++++++++++++++++++++++++++ client/android/session.go | 7 +- 6 files changed, 354 insertions(+), 17 deletions(-) create mode 100644 client/android/profile_state.go create mode 100644 client/android/profile_state_test.go diff --git a/client/android/client.go b/client/android/client.go index 59dcb1021..154bd8484 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -82,6 +82,8 @@ type Client struct { connectClient *internal.ConnectClient config *profilemanager.Config cacheDir string + // Identifies the running profile for the SSO login hint; see profile_state.go. + cfgPath string stateChangeMu sync.Mutex stateChangeSubID string @@ -102,11 +104,12 @@ type Client struct { extendCancel context.CancelFunc } -func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cc *internal.ConnectClient) { +func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cfgPath string, cc *internal.ConnectClient) { c.stateMu.Lock() defer c.stateMu.Unlock() c.config = cfg c.cacheDir = cacheDir + c.cfgPath = cfgPath c.connectClient = cc } @@ -116,6 +119,16 @@ func (c *Client) stateSnapshot() (*profilemanager.Config, string, *internal.Conn return c.config, c.cacheDir, c.connectClient } +// authSnapshot returns the config together with the path it was loaded from, in +// one lock: the path identifies the profile whose account email backs the login +// hint, so reading it separately could pair one profile's config with another's +// hint when a profile switch lands in between. +func (c *Client) authSnapshot() (*profilemanager.Config, string, *internal.ConnectClient) { + c.stateMu.RLock() + defer c.stateMu.RUnlock() + return c.config, c.cfgPath, c.connectClient +} + func (c *Client) getConnectClient() *internal.ConnectClient { c.stateMu.RLock() defer c.stateMu.RUnlock() @@ -168,7 +181,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid defer c.ctxCancel() c.ctxCancelLock.Unlock() - auth := NewAuthWithConfig(ctx, cfg) + auth := NewAuthWithConfig(ctx, cfg, cfgFile) err = auth.login(urlOpener, isAndroidTV) if err != nil { return err @@ -176,7 +189,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) - c.setState(cfg, cacheDir, connectClient) + c.setState(cfg, cacheDir, cfgFile, connectClient) // This path runs the interactive SSO flow, so reaching here means the peer // is authenticated again — release the latch Status() reports from. Clear // only once the fresh connect client is installed: until then Status() @@ -217,7 +230,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR // todo do not throw error in case of cancelled context ctx = internal.CtxInitState(ctx) connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) - c.setState(cfg, cacheDir, connectClient) + c.setState(cfg, cacheDir, cfgFile, connectClient) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } diff --git a/client/android/login.go b/client/android/login.go index 32ce28739..3f367b97f 100644 --- a/client/android/login.go +++ b/client/android/login.go @@ -4,6 +4,8 @@ import ( "context" "fmt" + log "github.com/sirupsen/logrus" + "github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/system" @@ -61,11 +63,14 @@ func NewAuth(cfgPath string, mgmURL string) (*Auth, error) { }, nil } -// NewAuthWithConfig instantiate Auth based on existing config -func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config) *Auth { +// NewAuthWithConfig instantiate Auth based on existing config. cfgPath is the +// file the config was loaded from; it identifies the profile whose account email +// backs the login_hint. +func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config, cfgPath string) *Auth { return &Auth{ - ctx: ctx, - config: config, + ctx: ctx, + config: config, + cfgPath: cfgPath, } } @@ -158,12 +163,14 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error { } jwtToken := "" + email := "" if needsLogin { tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV) if err != nil { return fmt.Errorf("interactive sso login failed: %v", err) } jwtToken = tokenInfo.GetTokenToUse() + email = tokenInfo.Email } err, _ = authClient.Login(a.ctx, "", jwtToken) @@ -171,17 +178,42 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error { return fmt.Errorf("login failed: %v", err) } + // Stored after Login, not before: a rejected token must not leave a hint + // pointing at an account that cannot be used. + if email != "" && a.cfgPath != "" { + if err := writeProfileEmail(a.cfgPath, email); err != nil { + log.Warnf("failed to store profile account email: %v", err) + } + } + go urlOpener.OnLoginSuccess() return nil } +// loginHintSetter is implemented by both concrete flows (PKCE and device code) +// but absent from the OAuthFlow interface, hence the assertion below — the same +// way internal/auth wires it in authenticateWithPKCEFlow. +type loginHintSetter interface { + SetLoginHint(hint string) +} + func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) { oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV) if err != nil { return nil, fmt.Errorf("failed to get OAuth flow: %v", err) } + // An empty hint is deliberate, not a fallback: a fresh or logged-out profile + // leaves the choice to the IdP, which is how accounts get switched. + if a.cfgPath != "" { + if hint := readProfileEmail(a.cfgPath); hint != "" { + if setter, ok := oAuthFlow.(loginHintSetter); ok { + setter.SetLoginHint(hint) + } + } + } + flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO()) if err != nil { return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err) diff --git a/client/android/profile_manager.go b/client/android/profile_manager.go index 9a051137c..3197124d7 100644 --- a/client/android/profile_manager.go +++ b/client/android/profile_manager.go @@ -13,18 +13,17 @@ import ( ) const ( - // Android-specific config filename (different from desktop default.json) - defaultConfigFilename = "netbird.cfg" - // Subdirectory for non-default profiles (must match Java Preferences.java) - profilesSubdir = "profiles" // Android uses a single user context per app (non-empty username required by ServiceManager) androidUsername = "android" ) // Profile represents a profile for gomobile type Profile struct { - ID string - Name string + ID string + Name string + // Email is the account this profile last logged in with, "" if it never + // completed an SSO login or was logged out. See profile_state.go. + Email string IsActive bool } @@ -101,6 +100,7 @@ func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) { profiles = append(profiles, &Profile{ ID: p.ID.String(), Name: p.Name, + Email: pm.profileEmail(p.ID.String()), IsActive: p.IsActive, }) } @@ -123,7 +123,22 @@ func (pm *ProfileManager) GetActiveProfile() (*Profile, error) { if err != nil { return nil, fmt.Errorf("failed to resolve active profile %q: %w", activeState.ID, err) } - return &Profile{ID: prof.ID.String(), Name: prof.Name, IsActive: true}, nil + return &Profile{ + ID: prof.ID.String(), + Name: prof.Name, + Email: pm.profileEmail(prof.ID.String()), + IsActive: true, + }, nil +} + +// profileEmail returns the account email recorded for a profile. Display-only, so +// an unresolvable path degrades to "" rather than an error. +func (pm *ProfileManager) profileEmail(id string) string { + configPath, err := pm.getProfileConfigPath(id) + if err != nil { + return "" + } + return readProfileEmail(configPath) } // SwitchProfile switches to a different profile @@ -185,6 +200,11 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return fmt.Errorf("failed to save config: %w", err) } + // Not fatal: a stale hint costs an account switch, not the logout itself. + if err := removeProfileEmail(configPath); err != nil { + log.Warnf("failed to clear stored account email for profile %s: %v", id, err) + } + log.Infof("logged out from profile: %s", id) return nil } diff --git a/client/android/profile_state.go b/client/android/profile_state.go new file mode 100644 index 000000000..3f0a09701 --- /dev/null +++ b/client/android/profile_state.go @@ -0,0 +1,108 @@ +package android + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/client/internal/profilemanager" + "github.com/netbirdio/netbird/util" +) + +const ( + // Android-specific config filename (different from desktop default.json) + defaultConfigFilename = "netbird.cfg" + // Subdirectory for non-default profiles (must match Java Preferences.java) + profilesSubdir = "profiles" + // profileAccountSuffix names the file holding the profile's account email. + // Deliberately not ".state.json", which desktop uses for the same data: + // there the email and the engine's state manager live in different + // directories, but on Android both resolve under files/, so sharing the name + // would have the two overwrite each other — the state manager rewrites the + // whole file from its own keys (see statemanager.Manager.PersistState), and + // this package's writer does the same in reverse. + profileAccountSuffix = ".account.json" +) + +// profileAccountPathFor derives the account file path from a profile's config +// path: netbird.cfg -> netbird.account.json, .json -> .account.json. +// +// Deriving from the config path rather than resolving the active profile keeps +// the write on the profile the login actually ran for: Auth.login runs in a +// goroutine, so the active profile can change under a flow already in flight. +func profileAccountPathFor(configPath string) (string, error) { + if configPath == "" { + return "", fmt.Errorf("empty config path") + } + + base := filepath.Base(configPath) + stem := strings.TrimSuffix(base, filepath.Ext(base)) + if stem == "" || stem == "." { + return "", fmt.Errorf("config path %q has no filename stem", configPath) + } + + return filepath.Join(filepath.Dir(configPath), stem+profileAccountSuffix), nil +} + +// readProfileEmail returns the account email stored for the profile whose config +// lives at configPath. A missing or unreadable file yields "", which leaves the +// account choice to the IdP. +func readProfileEmail(configPath string) string { + accountPath, err := profileAccountPathFor(configPath) + if err != nil { + log.Debugf("no profile account path for login hint: %v", err) + return "" + } + + var state profilemanager.ProfileState + if _, err := util.ReadJson(accountPath, &state); err != nil { + if !os.IsNotExist(err) { + log.Debugf("failed to read profile account for login hint: %v", err) + } + return "" + } + + return state.Email +} + +// writeProfileEmail records the account email for the profile whose config lives +// at configPath, so later logins can pass it as an OIDC login_hint. An empty +// email is ignored rather than blanking what is already stored. +func writeProfileEmail(configPath string, email string) error { + if email == "" { + return nil + } + + accountPath, err := profileAccountPathFor(configPath) + if err != nil { + return fmt.Errorf("resolve profile account path: %w", err) + } + + state := profilemanager.ProfileState{Email: email} + if err := util.WriteJsonWithRestrictedPermission(context.Background(), accountPath, state); err != nil { + return fmt.Errorf("write profile account: %w", err) + } + + return nil +} + +// removeProfileEmail drops the stored account email. Called on logout: while the +// email is on disk it goes out as a login_hint, which would steer the next login +// straight back into the account just logged out of. Mirrors the desktop UI's +// RemoveProfileState call. +func removeProfileEmail(configPath string) error { + accountPath, err := profileAccountPathFor(configPath) + if err != nil { + return fmt.Errorf("resolve profile account path: %w", err) + } + + if err := os.Remove(accountPath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("remove profile account: %w", err) + } + + return nil +} diff --git a/client/android/profile_state_test.go b/client/android/profile_state_test.go new file mode 100644 index 000000000..412435bc3 --- /dev/null +++ b/client/android/profile_state_test.go @@ -0,0 +1,161 @@ +package android + +import ( + "os" + "path/filepath" + "testing" +) + +func TestProfileAccountPathFor(t *testing.T) { + tests := []struct { + name string + configPath string + want string + wantErr bool + }{ + { + name: "default profile", + configPath: "/data/data/io.netbird.client/files/netbird.cfg", + want: "/data/data/io.netbird.client/files/netbird.account.json", + }, + { + name: "id profile", + configPath: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.json", + want: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.account.json", + }, + { + name: "legacy name-keyed profile is handled the same way", + configPath: "/data/data/io.netbird.client/files/profiles/work.json", + want: "/data/data/io.netbird.client/files/profiles/work.account.json", + }, + { + name: "empty path is rejected", + configPath: "", + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := profileAccountPathFor(tt.configPath) + if tt.wantErr { + if err == nil { + t.Fatalf("expected an error, got path %q", got) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != tt.want { + t.Errorf("got %q, want %q", got, tt.want) + } + }) + } +} + +func TestProfileAccountPathForDefaultDoesNotCollide(t *testing.T) { + root := "/data/data/io.netbird.client/files" + + defaultAccount, err := profileAccountPathFor(filepath.Join(root, defaultConfigFilename)) + if err != nil { + t.Fatalf("default profile: %v", err) + } + + idAccount, err := profileAccountPathFor(filepath.Join(root, profilesSubdir, "abc123.json")) + if err != nil { + t.Fatalf("id profile: %v", err) + } + + if defaultAccount == idAccount { + t.Fatalf("default and id profile share an account file: %q", defaultAccount) + } +} + +// The account file must never land on the engine state file: on Android both +// resolve under files/, and the state manager rewrites the whole file from its +// own keys, so sharing a path would have the two overwrite each other. The +// expected names here mirror ProfileManager.GetStateFilePath. +func TestProfileAccountPathAvoidsEngineStateFile(t *testing.T) { + root := "/data/data/io.netbird.client/files" + + cases := []struct { + configPath string + engineState string + }{ + { + configPath: filepath.Join(root, defaultConfigFilename), + engineState: filepath.Join(root, "state.json"), + }, + { + configPath: filepath.Join(root, profilesSubdir, "abc123.json"), + engineState: filepath.Join(root, profilesSubdir, "abc123.state.json"), + }, + } + + for _, c := range cases { + account, err := profileAccountPathFor(c.configPath) + if err != nil { + t.Fatalf("%s: %v", c.configPath, err) + } + if account == c.engineState { + t.Errorf("account file collides with the engine state file: %q", account) + } + } +} + +func TestWriteThenReadProfileEmail(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json") + if err := ensureDirFor(t, configPath); err != nil { + t.Fatalf("prepare dir: %v", err) + } + + if got := readProfileEmail(configPath); got != "" { + t.Errorf("expected no email before a login, got %q", got) + } + + const email = "user@example.com" + if err := writeProfileEmail(configPath, email); err != nil { + t.Fatalf("write: %v", err) + } + + if got := readProfileEmail(configPath); got != email { + t.Errorf("got %q, want %q", got, email) + } + + if err := removeProfileEmail(configPath); err != nil { + t.Fatalf("remove: %v", err) + } + if got := readProfileEmail(configPath); got != "" { + t.Errorf("expected no email after logout, got %q", got) + } + + // Logout may run on a never-logged-in profile, so a second remove must pass. + if err := removeProfileEmail(configPath); err != nil { + t.Fatalf("second remove should be a no-op: %v", err) + } +} + +func TestWriteProfileEmailIgnoresEmpty(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json") + if err := ensureDirFor(t, configPath); err != nil { + t.Fatalf("prepare dir: %v", err) + } + + const email = "user@example.com" + if err := writeProfileEmail(configPath, email); err != nil { + t.Fatalf("write: %v", err) + } + if err := writeProfileEmail(configPath, ""); err != nil { + t.Fatalf("write empty: %v", err) + } + + if got := readProfileEmail(configPath); got != email { + t.Errorf("empty write clobbered the stored email: got %q, want %q", got, email) + } +} + +func ensureDirFor(t *testing.T, path string) error { + t.Helper() + return os.MkdirAll(filepath.Dir(path), 0o700) +} diff --git a/client/android/session.go b/client/android/session.go index 961d52528..d5da09c93 100644 --- a/client/android/session.go +++ b/client/android/session.go @@ -278,7 +278,7 @@ func (c *Client) endExtend() { } func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isAndroidTV bool) error { - cfg, _, cc := c.stateSnapshot() + cfg, cfgPath, cc := c.authSnapshot() if cfg == nil || cc == nil { return fmt.Errorf("engine is not running") } @@ -293,7 +293,10 @@ func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isA } defer authClient.Close() - a := &Auth{ctx: ctx, config: cfg} + // Passing the config path makes the flow pick up the login_hint: an extend + // renews the session of the account already signed in, so it must not stop to + // offer a choice. + a := NewAuthWithConfig(ctx, cfg, cfgPath) tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV) if err != nil { return fmt.Errorf("interactive sso login failed: %v", err) From 530021aec6b5efd8d67e443727fd58aa47618f5a Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 3 Aug 2026 21:19:59 +0000 Subject: [PATCH 27/34] [client] Update wails to v3.0.0-beta.3 (#7038) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Update wails to v3.0.0-beta.3 ## 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) - [ ] 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). ## 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 * **Chores** * Updated the application framework dependency to a newer beta release for improved compatibility and stability. --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index 26a194958..44d9e908f 100644 --- a/go.mod +++ b/go.mod @@ -114,7 +114,7 @@ require ( github.com/ti-mo/conntrack v0.5.1 github.com/ti-mo/netfilter v0.5.2 github.com/vmihailenco/msgpack/v5 v5.4.1 - github.com/wailsapp/wails/v3 v3.0.0-alpha2.117 + github.com/wailsapp/wails/v3 v3.0.0-beta.3 github.com/yusufpapurcu/wmi v1.2.4 github.com/zcalusic/sysinfo v1.1.3 go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.67.0 diff --git a/go.sum b/go.sum index 58e30a580..762038e9e 100644 --- a/go.sum +++ b/go.sum @@ -660,8 +660,8 @@ github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IU github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok= github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g= github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds= -github.com/wailsapp/wails/v3 v3.0.0-alpha2.117 h1:udyjqPG3AIgkod5QDR/WblCkpV8R86BFPSrsWxSyt5Y= -github.com/wailsapp/wails/v3 v3.0.0-alpha2.117/go.mod h1:74WH2FScMsgucZvHHvv7eOefDXCm/CjuIxqhhZgPhKg= +github.com/wailsapp/wails/v3 v3.0.0-beta.3 h1:BrcZunEBVucncRx+xgkk9TzlXU4qc0ygJuEhKAAGaeA= +github.com/wailsapp/wails/v3 v3.0.0-beta.3/go.mod h1:BzATbK71VFikMMMCo434wAi0QcaI03P+xeaWgDvQvjw= github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU= github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA= github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM= From bc7a15ab71d9e9ed944fa3ec17952e77a696595d Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Tue, 4 Aug 2026 08:59:09 +0900 Subject: [PATCH 28/34] [management] Align agent-network API contracts for API clients (#7026) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Work on the Terraform provider (terraform-provider-netbird #177–#183) surfaced places where the agent-network API broke its own contracts or deviated from the conventions the rest of the management API follows, forcing client-side workarounds. Settings reads now follow the settings-endpoint convention: GET always answers with a JSON object. Before bootstrap it returns the defaults with an empty cluster/subdomain/endpoint (previously 200 with a JSON `null` body, while the spec said 404). The settings PUT can bootstrap the account by carrying a `cluster` — previously the row could only come into existence through the first provider create, and a settings-first setup was impossible; a differing cluster on a bootstrapped account is rejected instead of silently ignored. PUT remains full-state. The provider PUT schema promised omit-preserves semantics for several operator-editable fields that the handler never delivered (it builds the row from the request, like every other update handler). The schema wording now matches the shipped full-state behavior; only the api_key (secret) and session keys stay preserved by the manager. Identity headers are always present in provider responses so an explicitly cleared value round-trips as an empty string. The Go REST client gains the full agent-network surface (catalog, providers, policies, guardrails, budget rules, settings), including a shim translating the legacy 200+`null` settings body from older servers into an `IsNotFound` error. Note for reviewers: the dashboard special-cased the `null` settings body; it needs a small follow-up for the new defaults response (in progress). --- e2e/agentnetwork/management_test.go | 11 + e2e/agentnetwork/settings_bootstrap_test.go | 114 ++++ .../agentnetwork/handlers/handlers_test.go | 16 +- .../handlers/providers_handler_test.go | 49 ++ .../agentnetwork/handlers/settings_handler.go | 25 +- .../handlers/settings_handler_test.go | 137 +++++ .../internals/modules/agentnetwork/manager.go | 153 ++++-- .../modules/agentnetwork/types/provider.go | 36 +- .../agentnetwork/types/provider_test.go | 38 ++ .../modules/agentnetwork/types/settings.go | 48 +- .../agentnetwork_budgetrule_realstack_test.go | 15 +- shared/management/client/rest/agentnetwork.go | 381 ++++++++++++++ .../client/rest/agentnetwork_test.go | 497 ++++++++++++++++++ shared/management/client/rest/client.go | 5 + shared/management/http/api/openapi.yml | 44 +- shared/management/http/api/types.gen.go | 35 +- 16 files changed, 1464 insertions(+), 140 deletions(-) create mode 100644 e2e/agentnetwork/settings_bootstrap_test.go create mode 100644 management/internals/modules/agentnetwork/handlers/settings_handler_test.go create mode 100644 shared/management/client/rest/agentnetwork.go create mode 100644 shared/management/client/rest/agentnetwork_test.go diff --git a/e2e/agentnetwork/management_test.go b/e2e/agentnetwork/management_test.go index cfd03f63c..6962f9796 100644 --- a/e2e/agentnetwork/management_test.go +++ b/e2e/agentnetwork/management_test.go @@ -160,8 +160,19 @@ func TestSettingsRoundTrip(t *testing.T) { assert.Equal(t, before.Cluster, flipped.Cluster, "cluster must be immutable across updates") assert.Equal(t, before.Subdomain, flipped.Subdomain, "subdomain must be immutable across updates") + // A cluster different from the pinned one must be rejected; echoing the + // pinned one back is valid. + _, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr("attacker.cluster.invalid"), + EnableLogCollection: before.EnableLogCollection, + EnablePromptCollection: before.EnablePromptCollection, + RedactPii: before.RedactPii, + }) + requireClientError(t, err) + // Restore the original toggles. _, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr(before.Cluster), EnableLogCollection: before.EnableLogCollection, EnablePromptCollection: before.EnablePromptCollection, RedactPii: before.RedactPii, diff --git a/e2e/agentnetwork/settings_bootstrap_test.go b/e2e/agentnetwork/settings_bootstrap_test.go new file mode 100644 index 000000000..ea56f7064 --- /dev/null +++ b/e2e/agentnetwork/settings_bootstrap_test.go @@ -0,0 +1,114 @@ +//go:build e2e + +package agentnetwork + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/e2e/harness" + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// harnessStartFresh boots a dedicated combined server with its own fresh +// account and registers its teardown on t. +func harnessStartFresh(ctx context.Context, t *testing.T) (*harness.Combined, error) { + t.Helper() + fresh, err := harness.StartCombined(ctx) + if err != nil { + return nil, err + } + t.Cleanup(func() { _ = fresh.Terminate(context.Background()) }) + if _, err := fresh.Bootstrap(ctx); err != nil { + return nil, err + } + return fresh, nil +} + +// TestSettingsBootstrapViaPut covers the settings-first bootstrap path on an +// account that has never been bootstrapped: the GET reads as the defaults +// with an empty cluster/subdomain/endpoint, a PUT without a cluster has +// nothing to pin and fails, and a PUT carrying a cluster creates the row and +// pins it immutably. The shared srv cannot provide that starting state (any +// provider-creating test bootstraps it, and test order is deliberately not +// relied on), so this boots a dedicated combined server — the image is +// already built and cached by TestMain's StartCombined, so the extra cost is +// one container start. +func TestSettingsBootstrapViaPut(t *testing.T) { + ctx := context.Background() + + fresh, err := harnessStartFresh(ctx, t) + require.NoError(t, err, "start dedicated combined server") + + // Before agent-network bootstrap the settings read as the defaults, not + // as an error and not as a null body. + before, err := fresh.GetSettings(ctx) + require.NoError(t, err, "get settings on a fresh account must succeed") + assert.Empty(t, before.Cluster, "cluster must be empty before bootstrap") + assert.Empty(t, before.Subdomain, "subdomain must be empty before bootstrap") + assert.Empty(t, before.Endpoint, "endpoint must be empty before bootstrap, not a bare dot") + assert.True(t, before.EnableLogCollection, "defaults must show log collection on, matching bootstrap") + assert.False(t, before.EnablePromptCollection, "defaults must show prompt collection off") + + // A PUT without a cluster has nothing to pin the account to. + _, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + EnableLogCollection: true, + }) + requireClientError(t, err) + + // A PUT carrying a cluster bootstraps the account and applies the + // mutable fields from the same request. Every toggle is set away from + // its bootstrap default so each assertion can actually fail. + const cluster = "e2e.bootstrap.netbird.selfhosted" + bootstrapped, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr(cluster), + EnableLogCollection: false, + EnablePromptCollection: true, + RedactPii: true, + }) + require.NoError(t, err, "bootstrap settings via PUT must succeed") + assert.Equal(t, cluster, bootstrapped.Cluster, "cluster must be pinned from the request") + require.NotEmpty(t, bootstrapped.Subdomain, "subdomain must be assigned at bootstrap") + assert.Equal(t, bootstrapped.Subdomain+"."+cluster, bootstrapped.Endpoint, "endpoint must combine subdomain and cluster") + assert.False(t, bootstrapped.EnableLogCollection, "log collection from the bootstrap request must override the default") + assert.True(t, bootstrapped.EnablePromptCollection, "prompt collection from the bootstrap request must apply") + assert.True(t, bootstrapped.RedactPii, "redact toggle from the bootstrap request must apply") + + // The row is persisted: an independent read agrees on every field. + after, err := fresh.GetSettings(ctx) + require.NoError(t, err, "get settings after bootstrap must succeed") + assert.Equal(t, bootstrapped.Endpoint, after.Endpoint, "bootstrap must persist across reads") + assert.Equal(t, bootstrapped.EnableLogCollection, after.EnableLogCollection, "log collection must persist") + assert.Equal(t, bootstrapped.EnablePromptCollection, after.EnablePromptCollection, "prompt collection must persist") + assert.Equal(t, bootstrapped.RedactPii, after.RedactPii, "redact toggle must persist") + + // Once bootstrapped, later updates may omit the cluster entirely. + persisted, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + EnableLogCollection: true, + EnablePromptCollection: false, + RedactPii: true, + }) + require.NoError(t, err, "post-bootstrap update without cluster must succeed") + assert.Equal(t, cluster, persisted.Cluster, "omitted cluster must keep the pinned value") + assert.True(t, persisted.EnableLogCollection, "post-bootstrap toggle must apply") + assert.False(t, persisted.EnablePromptCollection, "post-bootstrap toggle must apply") + + // The cluster is immutable: a different value is rejected rather than + // silently ignored, and the rejected update must not disturb anything. + _, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{ + Cluster: ptr("other.cluster.invalid"), + EnableLogCollection: false, + }) + requireClientError(t, err) + + final, err := fresh.GetSettings(ctx) + require.NoError(t, err, "get settings after the rejected cluster change must succeed") + assert.Equal(t, persisted.Cluster, final.Cluster, "rejected update must not change the cluster") + assert.Equal(t, persisted.Endpoint, final.Endpoint, "rejected update must not change the endpoint") + assert.Equal(t, persisted.EnableLogCollection, final.EnableLogCollection, "rejected update must not apply its toggles") + assert.Equal(t, persisted.EnablePromptCollection, final.EnablePromptCollection, "rejected update must not apply its toggles") + assert.Equal(t, persisted.RedactPii, final.RedactPii, "rejected update must not apply its toggles") +} diff --git a/management/internals/modules/agentnetwork/handlers/handlers_test.go b/management/internals/modules/agentnetwork/handlers/handlers_test.go index 27ebea5dd..9d855c05d 100644 --- a/management/internals/modules/agentnetwork/handlers/handlers_test.go +++ b/management/internals/modules/agentnetwork/handlers/handlers_test.go @@ -17,6 +17,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/server/account" nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/store" @@ -61,10 +62,23 @@ func newAgentNetworkHandlerFixture(t *testing.T) *agentNetworkHandlerFixture { Return(true, context.Background(), nil). AnyTimes() - manager := agentnetwork.NewManager(st, perms, nil, nil) + // Swallow activity events so the mutation paths (create/update/delete) + // are exercisable through the HTTP layer. + accounts := account.NewMockManager(ctrl) + accounts.EXPECT(). + StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). + AnyTimes() + accounts.EXPECT(). + UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()). + AnyTimes() + + manager := agentnetwork.NewManager(st, perms, accounts, nil) h := &handler{manager: manager} router := mux.NewRouter() + router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST") + router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET") + router.HandleFunc("/agent-network/providers/{providerId}", h.updateProvider).Methods("PUT") h.addPolicyEndpoints(router) h.addConsumptionEndpoints(router) h.addBudgetRuleEndpoints(router) diff --git a/management/internals/modules/agentnetwork/handlers/providers_handler_test.go b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go index 649224c02..05024cde9 100644 --- a/management/internals/modules/agentnetwork/handlers/providers_handler_test.go +++ b/management/internals/modules/agentnetwork/handlers/providers_handler_test.go @@ -1,7 +1,9 @@ package handlers import ( + "encoding/json" "math" + nethttp "net/http" "testing" "github.com/stretchr/testify/assert" @@ -51,3 +53,50 @@ func TestValidate_ModelRates(t *testing.T) { assert.Error(t, validate(base(m), true), "case %q must be rejected", name) } } + +// TestProviderHandler_UpdateReplacesFullState pins the update contract shared +// with the other PUT endpoints: the request replaces the provider's mutable +// state, so optional fields absent from the JSON land as their zero values. +// The two exceptions are server-side: the api_key (a secret — omitted means +// "not rotated") and the session keypair, both preserved by the manager. The +// identity headers stay on the wire as explicit empty strings so a cleared +// value round-trips. +func TestProviderHandler_UpdateReplacesFullState(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + create := `{ + "provider_id": "openai_api", + "name": "openai", + "upstream_url": "https://api.openai.com", + "api_key": "sk-test", + "enabled": true, + "metadata_disabled": true, + "skip_tls_verification": true, + "extra_values": {"x-portkey-config": "pc-prod-3f2a"}, + "identity_header_user_id": "x-bf-dim-netbird_user_id", + "models": [{"id": "gpt-4o", "input_per_1k": 0.0025, "output_per_1k": 0.01}] + }` + rec := f.do(t, nethttp.MethodPost, "/agent-network/providers", create) + require.Equal(t, nethttp.StatusOK, rec.Code, "create must succeed: %s", rec.Body.String()) + + var created api.AgentNetworkProvider + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &created)) + + // Minimal update: only the required fields, no api_key. Everything + // optional must land as its zero value. + update := `{"provider_id": "openai_api", "name": "openai-renamed", "upstream_url": "https://api.openai.com", "enabled": true}` + rec = f.do(t, nethttp.MethodPut, "/agent-network/providers/"+created.Id, update) + require.Equal(t, nethttp.StatusOK, rec.Code, "update without api_key must succeed (key is preserved): %s", rec.Body.String()) + + var updated api.AgentNetworkProvider + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &updated)) + assert.Equal(t, "openai-renamed", updated.Name, "sent field must apply") + assert.True(t, updated.Enabled, "sent field must apply") + assert.False(t, updated.MetadataDisabled, "omitted metadata_disabled must land as false — PUT replaces the full state") + assert.False(t, updated.SkipTlsVerification, "omitted skip_tls_verification must land as false") + assert.Nil(t, updated.ExtraValues, "omitted extra_values must be cleared") + assert.Equal(t, "", updated.IdentityHeaderUserId, "omitted identity header must be cleared yet stay on the wire") + assert.Empty(t, updated.Models, "omitted models must be cleared") + assert.Contains(t, rec.Body.String(), `"identity_header_user_id":""`, + "cleared identity header must round-trip as an explicit empty string") +} diff --git a/management/internals/modules/agentnetwork/handlers/settings_handler.go b/management/internals/modules/agentnetwork/handlers/settings_handler.go index c65efad0f..171750838 100644 --- a/management/internals/modules/agentnetwork/handlers/settings_handler.go +++ b/management/internals/modules/agentnetwork/handlers/settings_handler.go @@ -2,7 +2,6 @@ package handlers import ( "encoding/json" - "errors" "net/http" "github.com/gorilla/mux" @@ -11,19 +10,20 @@ import ( nbcontext "github.com/netbirdio/netbird/management/server/context" "github.com/netbirdio/netbird/shared/management/http/api" "github.com/netbirdio/netbird/shared/management/http/util" - "github.com/netbirdio/netbird/shared/management/status" ) // addSettingsEndpoints registers the Agent Network settings routes. The -// settings row is bootstrapped server-side on first provider create; GET reads -// it and PUT updates the mutable collection toggles (cluster/subdomain stay -// immutable). +// settings row is bootstrapped server-side on first provider create or on the +// first PUT carrying a cluster; GET reads it and PUT applies a partial update +// of the mutable collection toggles (cluster/subdomain stay immutable). func (h *handler) addSettingsEndpoints(router *mux.Router) { router.HandleFunc("/agent-network/settings", h.getSettings).Methods("GET", "OPTIONS") router.HandleFunc("/agent-network/settings", h.updateSettings).Methods("PUT", "OPTIONS") } -// updateSettings applies the collection toggles to the account's settings row. +// updateSettings replaces the mutable settings fields on the account's row. +// A request carrying a cluster bootstraps the row when the account doesn't +// have one yet. func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) { userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) if err != nil { @@ -48,11 +48,9 @@ func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) { util.WriteJSONObject(r.Context(), w, updated.ToAPIResponse()) } -// getSettings returns the account's agent-network settings. The settings -// row is bootstrapped on first provider create, so freshly-onboarded -// accounts have nothing to read. Rather than 404-ing in that case (which -// the dashboard would have to special-case), return a JSON null with 200 -// so consumers can branch on the body alone. +// getSettings returns the account's agent-network settings. Accounts that +// haven't been bootstrapped yet read as the defaults with an empty cluster, +// subdomain and endpoint; the manager synthesises that view. func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) { userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) if err != nil { @@ -62,11 +60,6 @@ func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) { settings, err := h.manager.GetSettings(r.Context(), userAuth.AccountId, userAuth.UserId) if err != nil { - var sErr *status.Error - if errors.As(err, &sErr) && sErr.Type() == status.NotFound { - util.WriteJSONObject(r.Context(), w, nil) - return - } util.WriteError(r.Context(), err, w) return } diff --git a/management/internals/modules/agentnetwork/handlers/settings_handler_test.go b/management/internals/modules/agentnetwork/handlers/settings_handler_test.go new file mode 100644 index 000000000..636ec5b26 --- /dev/null +++ b/management/internals/modules/agentnetwork/handlers/settings_handler_test.go @@ -0,0 +1,137 @@ +package handlers + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// TestSettingsHandler_GetUnbootstrappedReturnsDefaults pins the settings-read +// convention shared with the account and DNS settings endpoints: settings +// always read as a JSON object. Before bootstrap that object carries the +// defaults with an empty cluster/subdomain/endpoint (the "not bootstrapped" +// signal) and no timestamps — never a 404 and never the legacy null body. +func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodGet, "/agent-network/settings", "") + require.Equal(t, http.StatusOK, rec.Code, + "unbootstrapped account must read as 200 with defaults: got %d body=%s", rec.Code, rec.Body.String()) + require.NotEqual(t, "null", trimSpace(rec.Body.String()), + "the legacy 200+null shape must not come back") + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.Empty(t, got.Cluster, "cluster must be empty until bootstrapped") + assert.Empty(t, got.Subdomain, "subdomain must be empty until bootstrapped") + assert.Empty(t, got.Endpoint, "endpoint must be empty until bootstrapped, not a bare dot") + assert.True(t, got.EnableLogCollection, "defaults must show log collection on, matching bootstrap") + assert.False(t, got.EnablePromptCollection, "defaults must show prompt collection off") + assert.False(t, got.RedactPii, "defaults must show redaction off") + require.NotNil(t, got.AccessLogRetentionDays) + assert.Equal(t, 30, *got.AccessLogRetentionDays, "defaults must show the bootstrap retention") + assert.Nil(t, got.CreatedAt, "no timestamps before a row exists") + assert.Nil(t, got.UpdatedAt, "no timestamps before a row exists") +} + +// TestSettingsHandler_PutBootstrapsWithCluster covers the settings-first +// bootstrap path: a PUT carrying a cluster on an unbootstrapped account +// creates the row (cluster pinned, subdomain assigned) and applies the +// mutable fields from the same request. +func TestSettingsHandler_PutBootstrapsWithCluster(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": false, "access_log_retention_days": 30}`) + require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String()) + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be pinned from the request") + assert.NotEmpty(t, got.Subdomain, "subdomain must be assigned at bootstrap") + assert.Equal(t, got.Subdomain+".eu.proxy.netbird.io", got.Endpoint, "endpoint must combine subdomain and cluster") + assert.True(t, got.EnableLogCollection, "toggle from the bootstrap request must apply") + assert.True(t, got.EnablePromptCollection, "toggle from the bootstrap request must apply") + require.NotNil(t, got.AccessLogRetentionDays) + assert.Equal(t, 30, *got.AccessLogRetentionDays, "retention from the bootstrap request must apply") + + // The row is now readable via GET. + rec = f.do(t, http.MethodGet, "/agent-network/settings", "") + require.Equal(t, http.StatusOK, rec.Code, "GET after bootstrap must succeed") +} + +// TestSettingsHandler_PutWithoutClusterOnUnbootstrapped pins that a PUT +// without a cluster cannot conjure a settings row out of nothing — there is +// no cluster to pin — and surfaces as 404 like the GET. +func TestSettingsHandler_PutWithoutClusterOnUnbootstrapped(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false}`) + assert.Equal(t, http.StatusNotFound, rec.Code, + "cluster-less PUT on an unbootstrapped account must 404: got %d body=%s", rec.Code, rec.Body.String()) + assert.Contains(t, rec.Body.String(), "cluster", + "the error must point the caller at the bootstrap paths: %s", rec.Body.String()) +} + +// TestSettingsHandler_PutReplacesMutableFields pins the update contract shared +// with the other PUT endpoints: the request replaces every mutable field, so a +// toggle absent from the JSON lands as its zero value rather than being +// preserved. Cluster and subdomain survive untouched. +func TestSettingsHandler_PutReplacesMutableFields(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": true, "access_log_retention_days": 14}`) + require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String()) + + var before api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before)) + + rec = f.do(t, http.MethodPut, "/agent-network/settings", + `{"enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`) + require.Equal(t, http.StatusOK, rec.Code, "update PUT must succeed: %s", rec.Body.String()) + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.True(t, got.EnableLogCollection, "sent toggle must apply") + assert.False(t, got.EnablePromptCollection, "sent toggle must apply") + assert.False(t, got.RedactPii, "sent toggle must apply") + require.NotNil(t, got.AccessLogRetentionDays) + assert.Equal(t, 0, *got.AccessLogRetentionDays, + "retention absent from the request must land as the zero value — PUT replaces all mutable fields") + assert.Equal(t, before.Cluster, got.Cluster, "cluster must survive updates untouched") + assert.Equal(t, before.Subdomain, got.Subdomain, "subdomain must survive updates untouched") +} + +// TestSettingsHandler_PutRejectsClusterChange pins cluster immutability: once +// assigned, a differing cluster is rejected as a validation error instead of +// being silently ignored, so callers never observe a value other than the one +// they sent. Echoing the assigned cluster back stays valid, which lets +// declarative clients send their full desired state idempotently. +func TestSettingsHandler_PutRejectsClusterChange(t *testing.T) { + f := newAgentNetworkHandlerFixture(t) + + rec := f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`) + require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String()) + + rec = f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "us.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`) + assert.Equal(t, http.StatusUnprocessableEntity, rec.Code, + "cluster change must be rejected as a validation error: got %d body=%s", rec.Code, rec.Body.String()) + + rec = f.do(t, http.MethodPut, "/agent-network/settings", + `{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": true}`) + require.Equal(t, http.StatusOK, rec.Code, "echoing the assigned cluster must stay valid: %s", rec.Body.String()) + + var got api.AgentNetworkSettings + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be unchanged") + assert.True(t, got.RedactPii, "toggle sent alongside the echoed cluster must apply") +} diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index 2687e4534..ba2c06826 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -207,7 +207,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide } if strings.TrimSpace(bootstrapCluster) != "" { - if _, err := m.bootstrapSettingsIfNeeded(ctx, provider.AccountID, bootstrapCluster); err != nil { + if _, err := m.bootstrapSettingsIfNeeded(ctx, m.store, provider.AccountID, bootstrapCluster); err != nil { // The provider create has already succeeded; logging the // bootstrap miss matches the plan's PoC behaviour. The synth // path treats a missing settings row as a no-op, and the next @@ -559,40 +559,83 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r return nil } -// UpdateSettings applies the mutable account-level settings — the collection -// toggles — onto the existing row. Cluster and Subdomain are immutable and are -// preserved from the persisted row regardless of the input. Because the -// collection toggles change the synthesised service config (prompt-capture -// gating, access-log emission), a reconcile is triggered so the proxy and peer -// network maps converge on the new state. +// UpdateSettings replaces the mutable account-level settings — the collection +// toggles and retention — on the account's row. When the account has no +// settings row yet, a non-empty settings.Cluster bootstraps one (same path as +// first provider create); without it the update fails with NotFound. On an +// existing row the cluster and subdomain are immutable: a differing +// settings.Cluster is rejected rather than silently ignored so callers never +// observe a value other than what they sent. Because the collection toggles +// change the synthesised service config (prompt-capture gating, access-log +// emission), a reconcile is triggered so the proxy and peer network maps +// converge on the new state. func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) { if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil { return nil, err } - existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID) + requestedCluster := strings.TrimSpace(settings.Cluster) + + // The row lock from LockingStrengthUpdate only holds for the duration of + // the surrounding transaction, so the read, the cluster-immutability + // check, and the save must share one — otherwise concurrent PUTs could + // interleave between them. + var updated *types.Settings + err := m.store.ExecuteInTransaction(ctx, func(tx store.Store) error { + existing, err := tx.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID) + switch { + case err == nil: + if requestedCluster != "" && requestedCluster != existing.Cluster { + return status.Errorf(status.InvalidArgument, "cluster is immutable once assigned (current: %s)", existing.Cluster) + } + case isNotFound(err): + if requestedCluster == "" { + return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; pass cluster to bootstrap them, or create a provider with bootstrap_cluster set") + } + // Bootstrapping pins the cluster and subdomain — a settings + // create on top of the update the caller already passed, matching + // the gate on the provider-create bootstrap path. + if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Create); err != nil { + return err + } + existing, err = m.bootstrapSettingsIfNeeded(ctx, tx, settings.AccountID, requestedCluster) + if err != nil { + return err + } + default: + return fmt.Errorf("get agent network settings: %w", err) + } + + existing.EnableLogCollection = settings.EnableLogCollection + existing.EnablePromptCollection = settings.EnablePromptCollection + existing.RedactPii = settings.RedactPii + existing.AccessLogRetentionDays = settings.AccessLogRetentionDays + existing.UpdatedAt = time.Now().UTC() + + if err := tx.SaveAgentNetworkSettings(ctx, existing); err != nil { + return fmt.Errorf("save agent network settings: %w", err) + } + updated = existing + return nil + }) if err != nil { - return nil, fmt.Errorf("get agent network settings: %w", err) - } - - existing.EnableLogCollection = settings.EnableLogCollection - existing.EnablePromptCollection = settings.EnablePromptCollection - existing.RedactPii = settings.RedactPii - existing.AccessLogRetentionDays = settings.AccessLogRetentionDays - existing.UpdatedAt = time.Now().UTC() - - if err := m.store.SaveAgentNetworkSettings(ctx, existing); err != nil { - return nil, fmt.Errorf("save agent network settings: %w", err) + return nil, err } m.accountManager.StoreEvent(ctx, userID, settings.AccountID, settings.AccountID, activity.AgentNetworkSettingsUpdated, map[string]any{ - "log_collection": existing.EnableLogCollection, - "prompt_collection": existing.EnablePromptCollection, - "redact_pii": existing.RedactPii, + "log_collection": updated.EnableLogCollection, + "prompt_collection": updated.EnablePromptCollection, + "redact_pii": updated.RedactPii, }) m.reconcile(ctx, settings.AccountID) - return existing, nil + return updated, nil +} + +// isNotFound reports whether err is a status.NotFound error. +func isNotFound(err error) bool { + var sErr *status.Error + return errors.As(err, &sErr) && sErr.Type() == status.NotFound } // validateProviderRefs ensures every destination provider id refers to a @@ -616,22 +659,25 @@ func (m *managerImpl) validateProviderRefs(ctx context.Context, accountID string return nil } -// GetSettings returns the agent-network settings row for the account. -// Returns the underlying status.NotFound when no row has been -// bootstrapped yet (i.e. the account has no providers). +// GetSettings returns the agent-network settings row for the account. When no +// row has been bootstrapped yet, the defaults are returned (without +// persisting) with cluster and subdomain empty — settings always read as an +// object, like the account and DNS settings endpoints. func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error) { if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Read); err != nil { return nil, err } - return m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + switch { + case err == nil: + return settings, nil + case isNotFound(err): + return types.DefaultSettings(accountID), nil + default: + return nil, err + } } -// bootstrapSettingsIfNeeded creates the per-account agent-network -// settings row when missing. The cluster comes from the create-time -// hint the dashboard sends (auto-picked from the active cluster list); -// the subdomain is picked from the curated wordlist avoiding -// collisions on the same cluster. Idempotent: if a row already exists -// it is returned untouched and the hint is ignored. // requireSettingsBootstrapPermission gates the one-time settings bootstrap a // first provider create performs. Pinning the account's cluster and subdomain // is a settings write, so it needs the settings permission on top of the @@ -641,14 +687,20 @@ func (m *managerImpl) requireSettingsBootstrapPermission(ctx context.Context, ac if err == nil { return nil } - var sErr *status.Error - if !errors.As(err, &sErr) || sErr.Type() != status.NotFound { + if !isNotFound(err) { return fmt.Errorf("get agent network settings: %w", err) } return m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Create) } -func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, providerCluster string) (*types.Settings, error) { +// bootstrapSettingsIfNeeded creates the per-account agent-network +// settings row when missing. The cluster comes from the create-time +// hint the dashboard sends (auto-picked from the active cluster list); +// the subdomain is picked from the curated wordlist avoiding +// collisions on the same cluster. Idempotent: if a row already exists +// it is returned untouched and the hint is ignored. st is the store to +// operate on — pass the transaction store when calling from within one. +func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, st store.Store, accountID, providerCluster string) (*types.Settings, error) { if accountID == "" { return nil, fmt.Errorf("bootstrap settings: account id is required") } @@ -656,16 +708,15 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, return nil, fmt.Errorf("bootstrap settings: provider cluster is required") } - existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) + existing, err := st.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID) if err == nil { return existing, nil } - var sErr *status.Error - if !errors.As(err, &sErr) || sErr.Type() != status.NotFound { + if !isNotFound(err) { return nil, fmt.Errorf("get agent network settings: %w", err) } - siblings, err := m.store.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster) + siblings, err := st.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster) if err != nil { return nil, fmt.Errorf("list agent network settings on cluster: %w", err) } @@ -684,18 +735,12 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, m.labelRngMu.Unlock() now := time.Now().UTC() - settings := &types.Settings{ - AccountID: accountID, - Cluster: providerCluster, - Subdomain: subdomain, - // Logs on by default; usage is collected regardless. Retention bounds - // how long full log rows are kept. - EnableLogCollection: true, - AccessLogRetentionDays: types.DefaultAccessLogRetentionDays, - CreatedAt: now, - UpdatedAt: now, - } - if err := m.store.SaveAgentNetworkSettings(ctx, settings); err != nil { + settings := types.DefaultSettings(accountID) + settings.Cluster = providerCluster + settings.Subdomain = subdomain + settings.CreatedAt = now + settings.UpdatedAt = now + if err := st.SaveAgentNetworkSettings(ctx, settings); err != nil { return nil, fmt.Errorf("save agent network settings: %w", err) } return settings, nil @@ -898,8 +943,8 @@ func (*mockManager) UpdateBudgetRule(_ context.Context, _ string, r *types.Accou func (*mockManager) DeleteBudgetRule(_ context.Context, _, _, _ string) error { return nil } -func (*mockManager) GetSettings(_ context.Context, _, _ string) (*types.Settings, error) { - return nil, status.Errorf(status.NotFound, "agent network settings not found") +func (*mockManager) GetSettings(_ context.Context, accountID, _ string) (*types.Settings, error) { + return types.DefaultSettings(accountID), nil } func (*mockManager) UpdateSettings(_ context.Context, _ string, s *types.Settings) (*types.Settings, error) { diff --git a/management/internals/modules/agentnetwork/types/provider.go b/management/internals/modules/agentnetwork/types/provider.go index 96242f45f..b9a194bf6 100644 --- a/management/internals/modules/agentnetwork/types/provider.go +++ b/management/internals/modules/agentnetwork/types/provider.go @@ -164,9 +164,7 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) { p.MetadataDisabled = *req.MetadataDisabled } // Identity-header overrides for catalogs flagged Customizable. - // nil pointer = "field omitted on the wire" → leave the stored - // value untouched (per the openapi description). Empty string is - // an explicit clear that disables stamping for this dimension. + // Empty or omitted disables stamping for this dimension. if req.IdentityHeaderUserId != nil { p.IdentityHeaderUserID = strings.TrimSpace(*req.IdentityHeaderUserId) } @@ -192,16 +190,20 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { created := p.CreatedAt updated := p.UpdatedAt resp := &api.AgentNetworkProvider{ - Id: p.ID, - ProviderId: p.ProviderID, - Name: p.Name, - UpstreamUrl: p.UpstreamURL, - Models: models, - Enabled: p.Enabled, - SkipTlsVerification: p.SkipTLSVerification, - MetadataDisabled: p.MetadataDisabled, - CreatedAt: &created, - UpdatedAt: &updated, + Id: p.ID, + ProviderId: p.ProviderID, + Name: p.Name, + UpstreamUrl: p.UpstreamURL, + Models: models, + // Always present on the wire so an explicitly cleared header + // round-trips as "" instead of vanishing from the response. + IdentityHeaderUserId: p.IdentityHeaderUserID, + IdentityHeaderGroups: p.IdentityHeaderGroups, + Enabled: p.Enabled, + SkipTlsVerification: p.SkipTLSVerification, + MetadataDisabled: p.MetadataDisabled, + CreatedAt: &created, + UpdatedAt: &updated, } if len(p.ExtraValues) > 0 { out := make(map[string]string, len(p.ExtraValues)) @@ -210,14 +212,6 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider { } resp.ExtraValues = &out } - if p.IdentityHeaderUserID != "" { - v := p.IdentityHeaderUserID - resp.IdentityHeaderUserId = &v - } - if p.IdentityHeaderGroups != "" { - v := p.IdentityHeaderGroups - resp.IdentityHeaderGroups = &v - } return resp } diff --git a/management/internals/modules/agentnetwork/types/provider_test.go b/management/internals/modules/agentnetwork/types/provider_test.go index f9756bb8b..fd553ece2 100644 --- a/management/internals/modules/agentnetwork/types/provider_test.go +++ b/management/internals/modules/agentnetwork/types/provider_test.go @@ -77,3 +77,41 @@ func TestProvider_MetadataDisabled_RoundTrip(t *testing.T) { assert.False(t, p.MetadataDisabled, "explicit false must clear metadata_disabled") assert.False(t, p.ToAPIResponse().MetadataDisabled, "response must reflect the cleared value") } + +// TestProvider_IdentityHeaders_AlwaysOnWire pins that the identity header +// fields are always present in the API response — an explicitly cleared +// ("") header must round-trip as "" rather than vanish, so API consumers +// (e.g. the Terraform provider) never observe a value other than the one +// they wrote. +func TestProvider_IdentityHeaders_AlwaysOnWire(t *testing.T) { + set := "x-bf-dim-netbird_user_id" + empty := "" + + base := func() *api.AgentNetworkProviderRequest { + return &api.AgentNetworkProviderRequest{ + ProviderId: "custom", + Name: "bifrost", + UpstreamUrl: "https://bifrost.internal", + } + } + + p := NewProvider("acc-1") + resp := p.ToAPIResponse() + assert.Equal(t, "", resp.IdentityHeaderUserId, "unset header must surface as empty string, not be omitted") + assert.Equal(t, "", resp.IdentityHeaderGroups, "unset header must surface as empty string, not be omitted") + + req := base() + req.IdentityHeaderUserId = &set + p.FromAPIRequest(req) + assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "configured header must round-trip") + + // Omitting the field preserves it. + p.FromAPIRequest(base()) + assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "omitted header must preserve the stored value") + + // An explicit "" clears it AND stays visible on the wire. + req = base() + req.IdentityHeaderUserId = &empty + p.FromAPIRequest(req) + assert.Equal(t, "", p.ToAPIResponse().IdentityHeaderUserId, "cleared header must round-trip as empty string") +} diff --git a/management/internals/modules/agentnetwork/types/settings.go b/management/internals/modules/agentnetwork/types/settings.go index d61d9deff..2c53877b5 100644 --- a/management/internals/modules/agentnetwork/types/settings.go +++ b/management/internals/modules/agentnetwork/types/settings.go @@ -1,6 +1,7 @@ package types import ( + "strings" "time" "github.com/netbirdio/netbird/shared/management/http/api" @@ -42,18 +43,34 @@ type Settings struct { // schema cohesive. func (Settings) TableName() string { return "agent_network_settings" } +// DefaultSettings returns the settings an account observes before its row is +// bootstrapped: log collection on with the default retention, everything else +// off, and no cluster/subdomain assigned yet. Bootstrap persists exactly these +// values plus the assigned cluster and subdomain, so the pre-bootstrap read +// and the freshly bootstrapped row agree. +func DefaultSettings(accountID string) *Settings { + return &Settings{ + AccountID: accountID, + EnableLogCollection: true, + AccessLogRetentionDays: DefaultAccessLogRetentionDays, + } +} + // Endpoint returns the bare hostname agents reach this account at: -// `.`. +// `.`. Empty until both halves are assigned at bootstrap. func (s *Settings) Endpoint() string { + if s.Cluster == "" || s.Subdomain == "" { + return "" + } return s.Subdomain + "." + s.Cluster } -// ToAPIResponse renders the settings as the API representation. +// ToAPIResponse renders the settings as the API representation. The +// timestamps are omitted while zero — a default (not yet bootstrapped) view +// has no persisted row to date. func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings { - created := s.CreatedAt - updated := s.UpdatedAt retention := s.AccessLogRetentionDays - return &api.AgentNetworkSettings{ + resp := &api.AgentNetworkSettings{ Cluster: s.Cluster, Subdomain: s.Subdomain, Endpoint: s.Endpoint(), @@ -61,14 +78,27 @@ func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings { EnablePromptCollection: s.EnablePromptCollection, RedactPii: s.RedactPii, AccessLogRetentionDays: &retention, - CreatedAt: &created, - UpdatedAt: &updated, } + if !s.CreatedAt.IsZero() { + created := s.CreatedAt + resp.CreatedAt = &created + } + if !s.UpdatedAt.IsZero() { + updated := s.UpdatedAt + resp.UpdatedAt = &updated + } + return resp } -// FromAPIRequest applies the mutable settings fields from the request. Cluster -// and Subdomain are immutable and intentionally not touched here. +// FromAPIRequest applies the request onto the receiver. The mutable +// collection fields are always replaced with the request values. Cluster +// participates only in bootstrap and the immutability check (see +// Manager.UpdateSettings); Subdomain is server-assigned and never taken +// from a request. func (s *Settings) FromAPIRequest(req *api.AgentNetworkSettingsRequest) { + if req.Cluster != nil { + s.Cluster = strings.TrimSpace(*req.Cluster) + } s.EnableLogCollection = req.EnableLogCollection s.EnablePromptCollection = req.EnablePromptCollection s.RedactPii = req.RedactPii diff --git a/management/server/agentnetwork_budgetrule_realstack_test.go b/management/server/agentnetwork_budgetrule_realstack_test.go index d17f2e26a..790285b4e 100644 --- a/management/server/agentnetwork_budgetrule_realstack_test.go +++ b/management/server/agentnetwork_budgetrule_realstack_test.go @@ -102,11 +102,20 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t require.NotEmpty(t, before.Subdomain, "subdomain pinned at bootstrap") assert.False(t, before.EnablePromptCollection, "prompt collection defaults off") - // Attempt to flip toggles AND smuggle a different cluster/subdomain — the - // immutable fields must be ignored. + // A cluster different from the one pinned at bootstrap must be rejected + // outright — never silently swapped or ignored. + _, err = mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{ + AccountID: accountID, + Cluster: "attacker.cluster", + EnableLogCollection: true, + }) + require.Error(t, err, "UpdateSettings with a mismatched cluster must fail") + + // Flipping the toggles works with the pinned cluster echoed back (and + // with it omitted); the subdomain is never taken from the request. updated, err := mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{ AccountID: accountID, - Cluster: "attacker.cluster", + Cluster: clusterAddr, Subdomain: "evil", EnableLogCollection: true, EnablePromptCollection: true, diff --git a/shared/management/client/rest/agentnetwork.go b/shared/management/client/rest/agentnetwork.go new file mode 100644 index 000000000..cee053d17 --- /dev/null +++ b/shared/management/client/rest/agentnetwork.go @@ -0,0 +1,381 @@ +package rest + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// AgentNetworkAPI APIs for the Agent Network (AI/LLM gateway), do not use directly +// see more: https://docs.netbird.io/api/resources/agent-network +type AgentNetworkAPI struct { + c *Client +} + +// ListCatalogProviders lists the catalog of supported upstream AI providers +// (openai_api, anthropic_api, bedrock_api, ...) with their default models and +// pricing, used to prefill provider create forms. +func (a *AgentNetworkAPI) ListCatalogProviders(ctx context.Context) ([]api.AgentNetworkCatalogProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/catalog/providers", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkCatalogProvider](resp) + return ret, err +} + +// ListProviders lists all Agent Network providers +func (a *AgentNetworkAPI) ListProviders(ctx context.Context) ([]api.AgentNetworkProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/providers", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkProvider](resp) + return ret, err +} + +// GetProvider gets Agent Network provider info +func (a *AgentNetworkAPI) GetProvider(ctx context.Context, providerID string) (*api.AgentNetworkProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/providers/"+providerID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// CreateProvider creates a new Agent Network provider. Set +// request.BootstrapCluster on the account's first provider to bootstrap the +// per-account gateway endpoint (alternatively bootstrap via UpdateSettings +// with a cluster). +func (a *AgentNetworkAPI) CreateProvider(ctx context.Context, request api.PostApiAgentNetworkProvidersJSONRequestBody) (*api.AgentNetworkProvider, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/providers", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// UpdateProvider updates an Agent Network provider. The request replaces the +// provider's mutable state; only an omitted api_key keeps the stored key +// (secrets are never required to round-trip). +func (a *AgentNetworkAPI) UpdateProvider(ctx context.Context, providerID string, request api.PutApiAgentNetworkProvidersProviderIdJSONRequestBody) (*api.AgentNetworkProvider, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/providers/"+providerID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// DeleteProvider deletes an Agent Network provider. Fails while any policy +// still references the provider — detach it first. +func (a *AgentNetworkAPI) DeleteProvider(ctx context.Context, providerID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/providers/"+providerID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListPolicies lists all Agent Network policies +func (a *AgentNetworkAPI) ListPolicies(ctx context.Context) ([]api.AgentNetworkPolicy, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/policies", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkPolicy](resp) + return ret, err +} + +// GetPolicy gets Agent Network policy info +func (a *AgentNetworkAPI) GetPolicy(ctx context.Context, policyID string) (*api.AgentNetworkPolicy, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/policies/"+policyID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// CreatePolicy creates a new Agent Network policy +func (a *AgentNetworkAPI) CreatePolicy(ctx context.Context, request api.PostApiAgentNetworkPoliciesJSONRequestBody) (*api.AgentNetworkPolicy, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/policies", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// UpdatePolicy updates an Agent Network policy +func (a *AgentNetworkAPI) UpdatePolicy(ctx context.Context, policyID string, request api.PutApiAgentNetworkPoliciesPolicyIdJSONRequestBody) (*api.AgentNetworkPolicy, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/policies/"+policyID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// DeletePolicy deletes an Agent Network policy +func (a *AgentNetworkAPI) DeletePolicy(ctx context.Context, policyID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/policies/"+policyID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListGuardrails lists all Agent Network guardrails +func (a *AgentNetworkAPI) ListGuardrails(ctx context.Context) ([]api.AgentNetworkGuardrail, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/guardrails", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkGuardrail](resp) + return ret, err +} + +// GetGuardrail gets Agent Network guardrail info +func (a *AgentNetworkAPI) GetGuardrail(ctx context.Context, guardrailID string) (*api.AgentNetworkGuardrail, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/guardrails/"+guardrailID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// CreateGuardrail creates a new Agent Network guardrail +func (a *AgentNetworkAPI) CreateGuardrail(ctx context.Context, request api.PostApiAgentNetworkGuardrailsJSONRequestBody) (*api.AgentNetworkGuardrail, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/guardrails", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// UpdateGuardrail updates an Agent Network guardrail +func (a *AgentNetworkAPI) UpdateGuardrail(ctx context.Context, guardrailID string, request api.PutApiAgentNetworkGuardrailsGuardrailIdJSONRequestBody) (*api.AgentNetworkGuardrail, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/guardrails/"+guardrailID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// DeleteGuardrail deletes an Agent Network guardrail +func (a *AgentNetworkAPI) DeleteGuardrail(ctx context.Context, guardrailID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/guardrails/"+guardrailID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListBudgetRules lists all account-level Agent Network budget rules +func (a *AgentNetworkAPI) ListBudgetRules(ctx context.Context) ([]api.AgentNetworkBudgetRule, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/budget-rules", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkBudgetRule](resp) + return ret, err +} + +// GetBudgetRule gets Agent Network budget rule info +func (a *AgentNetworkAPI) GetBudgetRule(ctx context.Context, ruleID string) (*api.AgentNetworkBudgetRule, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/budget-rules/"+ruleID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// CreateBudgetRule creates a new Agent Network budget rule +func (a *AgentNetworkAPI) CreateBudgetRule(ctx context.Context, request api.PostApiAgentNetworkBudgetRulesJSONRequestBody) (*api.AgentNetworkBudgetRule, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/budget-rules", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// UpdateBudgetRule updates an Agent Network budget rule +func (a *AgentNetworkAPI) UpdateBudgetRule(ctx context.Context, ruleID string, request api.PutApiAgentNetworkBudgetRulesRuleIdJSONRequestBody) (*api.AgentNetworkBudgetRule, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/budget-rules/"+ruleID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// DeleteBudgetRule deletes an Agent Network budget rule +func (a *AgentNetworkAPI) DeleteBudgetRule(ctx context.Context, ruleID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/budget-rules/"+ruleID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// GetSettings gets the account's Agent Network gateway settings (cluster, +// subdomain, endpoint, collection toggles). An account that has not been +// bootstrapped yet — via UpdateSettings with a cluster, or by creating the +// first provider with bootstrap_cluster set — reads as the defaults with an +// empty Cluster, Subdomain and Endpoint. Management servers prior to that +// contract answered 200 with a JSON null body instead; that legacy shape is +// translated to an APIError matchable via IsNotFound rather than fabricating +// defaults the server never stated. +func (a *AgentNetworkAPI) GetSettings(ctx context.Context) (*api.AgentNetworkSettings, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/settings", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, err + } + if trimmed := bytes.TrimSpace(body); len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil, &APIError{StatusCode: http.StatusNotFound, Message: "agent network settings not found"} + } + var ret api.AgentNetworkSettings + if err := json.Unmarshal(body, &ret); err != nil { + return nil, err + } + return &ret, nil +} + +// UpdateSettings updates the account's Agent Network settings; the request +// replaces every mutable field (collection toggles and retention). Setting +// request.Cluster bootstraps the settings row when the account does not have +// one yet; on a bootstrapped account it must match the assigned cluster (or +// be nil) and any other value is rejected — the cluster is immutable. +func (a *AgentNetworkAPI) UpdateSettings(ctx context.Context, request api.PutApiAgentNetworkSettingsJSONRequestBody) (*api.AgentNetworkSettings, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/settings", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkSettings](resp) + return &ret, err +} diff --git a/shared/management/client/rest/agentnetwork_test.go b/shared/management/client/rest/agentnetwork_test.go new file mode 100644 index 000000000..053859125 --- /dev/null +++ b/shared/management/client/rest/agentnetwork_test.go @@ -0,0 +1,497 @@ +//go:build integration + +package rest_test + +import ( + "context" + "encoding/json" + "io" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/client/rest" + "github.com/netbirdio/netbird/shared/management/http/api" + "github.com/netbirdio/netbird/shared/management/http/util" +) + +var ( + testAgentNetworkProvider = api.AgentNetworkProvider{ + Id: "ainp_test", + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + Models: []api.AgentNetworkProviderModel{}, + Enabled: true, + } + + testAgentNetworkPolicy = api.AgentNetworkPolicy{ + Id: "ainpol_test", + Name: "Engineering → OpenAI", + Enabled: true, + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + } + + testAgentNetworkGuardrail = api.AgentNetworkGuardrail{ + Id: "aingr_test", + Name: "No secrets", + } + + testAgentNetworkBudgetRule = api.AgentNetworkBudgetRule{ + Id: "ainbud_test", + Name: "Org monthly ceiling", + Enabled: true, + } + + testAgentNetworkSettings = api.AgentNetworkSettings{ + Cluster: "eu.proxy.netbird.io", + Subdomain: "violet", + Endpoint: "violet.eu.proxy.netbird.io", + EnableLogCollection: true, + AccessLogRetentionDays: ptr(30), + } +) + +func TestAgentNetwork_ListCatalogProviders_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/catalog/providers", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkCatalogProvider{{Id: "openai_api", Name: "OpenAI"}}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListCatalogProviders(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, "openai_api", ret[0].Id) + }) +} + +func TestAgentNetwork_ListProviders_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkProvider{testAgentNetworkProvider}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListProviders(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkProvider, ret[0]) + }) +} + +func TestAgentNetwork_GetProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "GET", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetProvider(context.Background(), "ainp_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_GetProvider_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "not found", Code: 404}) + w.WriteHeader(404) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.GetProvider(context.Background(), "ainp_test") + require.Error(t, err) + assert.True(t, rest.IsNotFound(err), "a 404 must be matchable via IsNotFound") + }) +} + +func TestAgentNetwork_CreateProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + var req api.PostApiAgentNetworkProvidersJSONRequestBody + require.NoError(t, json.Unmarshal(reqBytes, &req)) + assert.Equal(t, "OpenAI", req.Name) + require.NotNil(t, req.BootstrapCluster) + assert.Equal(t, "eu.proxy.netbird.io", *req.BootstrapCluster) + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateProvider(context.Background(), api.PostApiAgentNetworkProvidersJSONRequestBody{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + ApiKey: ptr("sk-test"), + BootstrapCluster: ptr("eu.proxy.netbird.io"), + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_UpdateProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + // Omitted optional fields must be absent from the wire (not + // zero-valued) so the server-side merge preserves them. + assert.NotContains(t, string(reqBytes), "api_key") + assert.NotContains(t, string(reqBytes), "models") + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateProvider(context.Background(), "ainp_test", api.PutApiAgentNetworkProvidersProviderIdJSONRequestBody{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_DeleteProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteProvider(context.Background(), "ainp_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListPolicies_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkPolicy{testAgentNetworkPolicy}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListPolicies(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkPolicy, ret[0]) + }) +} + +func TestAgentNetwork_GetPolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetPolicy(context.Background(), "ainpol_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_CreatePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreatePolicy(context.Background(), api.PostApiAgentNetworkPoliciesJSONRequestBody{ + Name: "Engineering → OpenAI", + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_UpdatePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdatePolicy(context.Background(), "ainpol_test", api.PutApiAgentNetworkPoliciesPolicyIdJSONRequestBody{ + Name: "Engineering → OpenAI", + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_DeletePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeletePolicy(context.Background(), "ainpol_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListGuardrails_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkGuardrail{testAgentNetworkGuardrail}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListGuardrails(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkGuardrail, ret[0]) + }) +} + +func TestAgentNetwork_GetGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetGuardrail(context.Background(), "aingr_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_CreateGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateGuardrail(context.Background(), api.PostApiAgentNetworkGuardrailsJSONRequestBody{ + Name: "No secrets", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_UpdateGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateGuardrail(context.Background(), "aingr_test", api.PutApiAgentNetworkGuardrailsGuardrailIdJSONRequestBody{ + Name: "No secrets", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_DeleteGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteGuardrail(context.Background(), "aingr_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListBudgetRules_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkBudgetRule{testAgentNetworkBudgetRule}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListBudgetRules(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkBudgetRule, ret[0]) + }) +} + +func TestAgentNetwork_GetBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetBudgetRule(context.Background(), "ainbud_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_CreateBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateBudgetRule(context.Background(), api.PostApiAgentNetworkBudgetRulesJSONRequestBody{ + Name: "Org monthly ceiling", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_UpdateBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateBudgetRule(context.Background(), "ainbud_test", api.PutApiAgentNetworkBudgetRulesRuleIdJSONRequestBody{ + Name: "Org monthly ceiling", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_DeleteBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteBudgetRule(context.Background(), "ainbud_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_GetSettings_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkSettings) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkSettings, *ret) + }) +} + +// TestAgentNetwork_GetSettings_UnbootstrappedDefaults pins the settings-read +// contract: an unbootstrapped account answers 200 with the defaults and empty +// cluster/subdomain/endpoint, which the client passes through untouched. +func TestAgentNetwork_GetSettings_UnbootstrappedDefaults(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(api.AgentNetworkSettings{ + EnableLogCollection: true, + AccessLogRetentionDays: ptr(30), + }) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.NoError(t, err) + assert.Empty(t, ret.Endpoint, "empty endpoint is the not-bootstrapped signal") + assert.True(t, ret.EnableLogCollection, "defaults must pass through") + }) +} + +func TestAgentNetwork_GetSettings_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "no", Code: 403}) + w.WriteHeader(403) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.GetSettings(context.Background()) + require.Error(t, err) + assert.Equal(t, "no", err.Error()) + }) +} + +// TestAgentNetwork_GetSettings_LegacyNullBody pins the compatibility shim for +// management servers that answered 200 with a JSON null body before the +// defaults contract: the client translates that shape into an IsNotFound +// error instead of returning a bogus zero-valued settings object or +// fabricating defaults the server never stated. +func TestAgentNetwork_GetSettings_LegacyNullBody(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + _, err := w.Write([]byte("null")) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.Error(t, err) + assert.Nil(t, ret) + assert.True(t, rest.IsNotFound(err), "the legacy 200+null shape must surface as IsNotFound") + }) +} + +func TestAgentNetwork_UpdateSettings_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + var req api.PutApiAgentNetworkSettingsJSONRequestBody + require.NoError(t, json.Unmarshal(reqBytes, &req)) + require.NotNil(t, req.Cluster, "bootstrap cluster must be on the wire") + assert.Equal(t, "eu.proxy.netbird.io", *req.Cluster) + assert.True(t, req.EnableLogCollection) + retBytes, _ := json.Marshal(testAgentNetworkSettings) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{ + Cluster: ptr("eu.proxy.netbird.io"), + EnableLogCollection: true, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkSettings, *ret) + }) +} + +func TestAgentNetwork_UpdateSettings_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "cluster is immutable once assigned (current: eu.proxy.netbird.io)", Code: 422}) + w.WriteHeader(422) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{ + Cluster: ptr("us.proxy.netbird.io"), + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "immutable") + }) +} diff --git a/shared/management/client/rest/client.go b/shared/management/client/rest/client.go index 43312b9e6..6154a6637 100644 --- a/shared/management/client/rest/client.go +++ b/shared/management/client/rest/client.go @@ -147,6 +147,10 @@ type Client struct { // ReverseProxyTokens account-scoped proxy access tokens used to register // self-hosted (bring-your-own-proxy) `netbird proxy` instances. ReverseProxyTokens *ReverseProxyTokensAPI + + // AgentNetwork NetBird Agent Network (AI/LLM gateway) APIs: catalog, + // providers, policies, guardrails, budget rules and account settings. + AgentNetwork *AgentNetworkAPI } // New initialize new Client instance using PAT token @@ -209,6 +213,7 @@ func (c *Client) initialize() { c.ReverseProxyClusters = &ReverseProxyClustersAPI{c} c.ReverseProxyDomains = &ReverseProxyDomainsAPI{c} c.ReverseProxyTokens = &ReverseProxyTokensAPI{c} + c.AgentNetwork = &AgentNetworkAPI{c} } // NewRequest creates and executes new management API request diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 8ad3d932c..551a60e2a 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -5149,12 +5149,12 @@ components: identity_header_user_id: type: string description: | - Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). + Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). example: "x-bf-dim-netbird_user_id" identity_header_groups: type: string description: | - Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. + Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. example: "x-bf-dim-netbird_groups" enabled: type: boolean @@ -5186,6 +5186,8 @@ components: - name - upstream_url - models + - identity_header_user_id + - identity_header_groups - enabled - skip_tls_verification - metadata_disabled @@ -5222,7 +5224,7 @@ components: extra_values: type: object description: | - Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). When present on a request, the whole map replaces the stored values. Empty strings drop the corresponding key. + Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). The request's map replaces the stored values; empty strings drop the corresponding key. additionalProperties: type: string example: @@ -5230,12 +5232,12 @@ components: identity_header_user_id: type: string description: | - Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. When omitted on a request, the stored value is left unchanged; pass an empty string explicitly to clear it (which disables stamping for this dimension). + Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. Empty or omitted disables stamping for this dimension. example: "x-bf-dim-netbird_user_id" identity_header_groups: type: string description: | - Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same omit / empty semantics as `identity_header_user_id`. + Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same semantics as `identity_header_user_id`. example: "x-bf-dim-netbird_groups" enabled: type: boolean @@ -5243,11 +5245,11 @@ components: example: true skip_tls_verification: type: boolean - description: Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. When omitted on update, the stored value is left unchanged. + description: Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. example: false metadata_disabled: type: boolean - description: Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). When omitted on update, the stored value is left unchanged. + description: Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). example: false required: - provider_id @@ -6191,19 +6193,19 @@ components: - cache_cost_usd AgentNetworkSettings: type: object - description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are auto-assigned on first provider create and immutable thereafter. + description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are assigned at bootstrap and immutable thereafter. Before bootstrap the account reads as the default values with empty cluster, subdomain and endpoint. properties: cluster: type: string - description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. + description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. Empty until the account is bootstrapped. example: "eu.proxy.netbird.io" subdomain: type: string - description: Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. + description: Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. Empty until the account is bootstrapped. example: "violet" endpoint: type: string - description: Bare hostname agents call for this account, computed as `.`. + description: Bare hostname agents call for this account, computed as `.`. Empty until the account is bootstrapped. example: "violet.eu.proxy.netbird.io" enable_log_collection: type: boolean @@ -6224,13 +6226,13 @@ components: created_at: type: string format: date-time - description: Timestamp when the settings row was created. + description: Timestamp when the settings row was created. Absent until the account is bootstrapped. readOnly: true example: "2026-04-26T10:30:00Z" updated_at: type: string format: date-time - description: Timestamp when the settings row was last updated. + description: Timestamp when the settings row was last updated. Absent until the account is bootstrapped. readOnly: true example: "2026-04-26T10:30:00Z" required: @@ -6240,12 +6242,14 @@ components: - enable_log_collection - enable_prompt_collection - redact_pii - - created_at - - updated_at AgentNetworkSettingsRequest: type: object - description: Mutable account-level Agent Network settings. Cluster and subdomain are immutable and not accepted here. + description: Account-level Agent Network settings update. The request replaces every mutable field. `cluster` additionally bootstraps the per-account settings row when the account does not have one yet; the subdomain is always server-assigned. properties: + cluster: + type: string + description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. When the account has no settings row yet, providing it bootstraps the row (assigning the subdomain that forms the agent endpoint). The cluster is immutable once assigned — later updates must omit it or send the assigned value; any other value is rejected. + example: "eu.proxy.netbird.io" enable_log_collection: type: boolean description: Whether per-request access-log entries are collected for this account's agent-network traffic. @@ -13690,7 +13694,7 @@ paths: /api/agent-network/settings: get: summary: Retrieve Agent Network settings - description: Returns the per-account Agent Network gateway settings (cluster, subdomain, endpoint). Returns 404 when no provider has been created yet — settings are lazily bootstrapped on first provider create. + description: Returns the per-account Agent Network gateway settings (cluster, subdomain, endpoint). Before the account is bootstrapped — on first provider create (`bootstrap_cluster`) or via PUT with `cluster` — the response carries the default values with empty cluster, subdomain and endpoint. tags: [ Agent Network ] security: - BearerAuth: [ ] @@ -13706,13 +13710,11 @@ paths: "$ref": "#/components/responses/requires_authentication" '403': "$ref": "#/components/responses/forbidden" - '404': - "$ref": "#/components/responses/not_found" '500': "$ref": "#/components/responses/internal_error" put: summary: Update Agent Network settings - description: Updates the mutable account-level Agent Network settings (collection toggles). Cluster and subdomain are immutable and ignored if sent. Returns 404 when settings have not been bootstrapped (no provider created yet). + description: Updates the account-level Agent Network settings; the request replaces every mutable field (collection toggles and retention). When the account has no settings row yet, providing `cluster` bootstraps it (assigning the subdomain that forms the agent endpoint); without `cluster` the request returns 404. Sending a `cluster` different from the assigned one is rejected (the cluster is immutable once assigned). The subdomain is always server-assigned and immutable. tags: [ Agent Network ] security: - BearerAuth: [ ] @@ -13738,6 +13740,8 @@ paths: "$ref": "#/components/responses/forbidden" '404': "$ref": "#/components/responses/not_found" + '422': + "$ref": "#/components/responses/validation_failed" '500': "$ref": "#/components/responses/internal_error" /api/agent-network/budget-rules: diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index 87dd9ccfc..dec644114 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -2275,11 +2275,11 @@ type AgentNetworkProvider struct { // Id Provider ID Id string `json:"id"` - // IdentityHeaderGroups Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. - IdentityHeaderGroups *string `json:"identity_header_groups,omitempty"` + // IdentityHeaderGroups Wire header name the proxy stamps with the caller's NetBird groups as a comma-separated list (sorted) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Same per-catalog semantics as `identity_header_user_id`. + IdentityHeaderGroups string `json:"identity_header_groups"` - // IdentityHeaderUserId Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). - IdentityHeaderUserId *string `json:"identity_header_user_id,omitempty"` + // IdentityHeaderUserId Wire header name the proxy stamps with the caller's display identity (user email or peer name) when the catalog entry's HeaderPair is `customizable`. Always present in responses; empty disables stamping for this dimension. Ignored when the catalog entry has a fixed HeaderPair (e.g. LiteLLM, Portkey). Used today by Bifrost: typical values are `x-bf-lh-netbird_user_id` (always-on log metadata) or `x-bf-dim-netbird_user_id` (Prometheus / OTEL — requires the label to be pre-declared in the gateway's `client.prometheus_labels` config). + IdentityHeaderUserId string `json:"identity_header_user_id"` // MetadataDisabled Whether identity metadata injection is disabled for this provider. When enabled (the default), the proxy stamps the caller's user and authorizing group onto upstream requests as provider-specific metadata (e.g. AWS Bedrock's X-Amzn-Bedrock-Request-Metadata header). Set true to suppress it. MetadataDisabled bool `json:"metadata_disabled"` @@ -2335,16 +2335,16 @@ type AgentNetworkProviderRequest struct { // Enabled Whether the provider is enabled. Defaults to true on create. Enabled *bool `json:"enabled,omitempty"` - // ExtraValues Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). When present on a request, the whole map replaces the stored values. Empty strings drop the corresponding key. + // ExtraValues Operator-typed values for catalog-declared extra headers (see AgentNetworkProvider.extra_values). The request's map replaces the stored values; empty strings drop the corresponding key. ExtraValues *map[string]string `json:"extra_values,omitempty"` - // IdentityHeaderGroups Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same omit / empty semantics as `identity_header_user_id`. + // IdentityHeaderGroups Wire header name for the caller's groups CSV. See AgentNetworkProvider.identity_header_groups. Same semantics as `identity_header_user_id`. IdentityHeaderGroups *string `json:"identity_header_groups,omitempty"` - // IdentityHeaderUserId Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. When omitted on a request, the stored value is left unchanged; pass an empty string explicitly to clear it (which disables stamping for this dimension). + // IdentityHeaderUserId Wire header name for the caller's display identity. See AgentNetworkProvider.identity_header_user_id. Empty or omitted disables stamping for this dimension. IdentityHeaderUserId *string `json:"identity_header_user_id,omitempty"` - // MetadataDisabled Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). When omitted on update, the stored value is left unchanged. + // MetadataDisabled Disable identity metadata injection (the caller's user + authorizing group) for this provider. Defaults to false (metadata is injected). MetadataDisabled *bool `json:"metadata_disabled,omitempty"` // Models Models exposed through this endpoint, with the operator's per-1k input/output prices. Empty means all catalog models are allowed at catalog prices. @@ -2356,22 +2356,22 @@ type AgentNetworkProviderRequest struct { // ProviderId Catalog identifier for the upstream AI provider (e.g. openai_api, anthropic_api, azure_openai_api, bedrock_api, vertex_ai_api, mistral_api, custom). ProviderId string `json:"provider_id"` - // SkipTlsVerification Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. When omitted on update, the stored value is left unchanged. + // SkipTlsVerification Skip upstream TLS certificate verification when the proxy dials this provider's URL. For self-hosted / internal gateways behind a private or self-signed certificate. Defaults to false. SkipTlsVerification *bool `json:"skip_tls_verification,omitempty"` // UpstreamUrl Full upstream URL (with scheme) that NetBird forwards traffic to. UpstreamUrl string `json:"upstream_url"` } -// AgentNetworkSettings Per-account Agent Network gateway settings. One row per account; cluster and subdomain are auto-assigned on first provider create and immutable thereafter. +// AgentNetworkSettings Per-account Agent Network gateway settings. One row per account; cluster and subdomain are assigned at bootstrap and immutable thereafter. Before bootstrap the account reads as the default values with empty cluster, subdomain and endpoint. type AgentNetworkSettings struct { // AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. Usage records are retained independently. AccessLogRetentionDays *int `json:"access_log_retention_days,omitempty"` - // Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. + // Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. Empty until the account is bootstrapped. Cluster string `json:"cluster"` - // CreatedAt Timestamp when the settings row was created. + // CreatedAt Timestamp when the settings row was created. Absent until the account is bootstrapped. CreatedAt *time.Time `json:"created_at,omitempty"` // EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic. @@ -2380,24 +2380,27 @@ type AgentNetworkSettings struct { // EnablePromptCollection Master switch for request/response prompt capture. Capture runs only when this is on AND a policy guardrail also enables it. EnablePromptCollection bool `json:"enable_prompt_collection"` - // Endpoint Bare hostname agents call for this account, computed as `.`. + // Endpoint Bare hostname agents call for this account, computed as `.`. Empty until the account is bootstrapped. Endpoint string `json:"endpoint"` // RedactPii Whether captured prompts have PII redacted. Effective redaction is the OR of this and any policy guardrail's redact setting. RedactPii bool `json:"redact_pii"` - // Subdomain Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. + // Subdomain Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. Empty until the account is bootstrapped. Subdomain string `json:"subdomain"` - // UpdatedAt Timestamp when the settings row was last updated. + // UpdatedAt Timestamp when the settings row was last updated. Absent until the account is bootstrapped. UpdatedAt *time.Time `json:"updated_at,omitempty"` } -// AgentNetworkSettingsRequest Mutable account-level Agent Network settings. Cluster and subdomain are immutable and not accepted here. +// AgentNetworkSettingsRequest Account-level Agent Network settings update. The request replaces every mutable field. `cluster` additionally bootstraps the per-account settings row when the account does not have one yet; the subdomain is always server-assigned. type AgentNetworkSettingsRequest struct { // AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. AccessLogRetentionDays *int `json:"access_log_retention_days,omitempty"` + // Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. When the account has no settings row yet, providing it bootstraps the row (assigning the subdomain that forms the agent endpoint). The cluster is immutable once assigned — later updates must omit it or send the assigned value; any other value is rejected. + Cluster *string `json:"cluster,omitempty"` + // EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic. EnableLogCollection bool `json:"enable_log_collection"` From 78c1c2fc32a9cc55f7fd1b81f5ea33e9d2873d11 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 4 Aug 2026 12:27:47 +0000 Subject: [PATCH 29/34] [client] Probe the daemon login with IsLoginRequired (#7052) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes Probe the daemon login with IsLoginRequired The Login probe attemptLogin(ctx, "", "") on an unregistered peer ends in registerPeer with no setup key and no JWT, which fails locally with InvalidArgument before reaching Management. Since #6983 classified that as StatusLoginFailed and returned early, every setup-key enrolment and every expired-session SSO re-login aborted before using its credentials, breaking all netbird-cloud e2e runs from commit e90be36cd. IsLoginRequired asks the question the probe actually means - is the peer's key alone still accepted - and reports Management's refusal as a decision (needsLogin) instead of an error, the same pattern foregroundLogin, Android and iOS already use. ## 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) - [ ] 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). ## 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 * **Bug Fixes** * Improved login handling when Management connectivity checks fail. * Prevented unnecessary SSO prompts for already-authenticated sessions. * Preserved setup-key login behavior while ensuring authentication attempts proceed correctly. * Login failures now return a clear failure status when authentication state cannot be verified. --- client/server/login_outcome_test.go | 39 ++++++++++++++++++++++------- client/server/server.go | 39 +++++++++++++++++++++-------- 2 files changed, 58 insertions(+), 20 deletions(-) diff --git a/client/server/login_outcome_test.go b/client/server/login_outcome_test.go index d3b1b7c52..7ebf04f92 100644 --- a/client/server/login_outcome_test.go +++ b/client/server/login_outcome_test.go @@ -8,8 +8,6 @@ import ( "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/proto" @@ -27,9 +25,9 @@ func TestLogin_ManagementUnreachableIsReturnedInsteadOfDemandingSSO(t *testing.T unreachable := errors.New("create connection: dial context: context deadline exceeded") attempts := 0 - s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) { + s.isLoginRequiredFn = func(context.Context) (bool, error) { attempts++ - return internal.StatusLoginFailed, unreachable + return false, unreachable } resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) @@ -55,15 +53,12 @@ func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) { s.rootCtx = internal.CtxInitState(context.Background()) breakProfilePrivateKey(t, cfgPath) - refused := gstatus.Error(codes.PermissionDenied, "peer is not registered") - s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) { - return internal.StatusNeedsLogin, refused + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return true, nil } _, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username}) require.Error(t, err) - require.NotErrorIs(t, err, refused, - "the refusal was handed back to the caller instead of starting the SSO flow") status, stateErr := internal.CtxGetState(s.rootCtx).Status() require.NoError(t, stateErr) @@ -71,6 +66,32 @@ func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) { "the SSO flow setup was never reached with the broken key") } +func TestLogin_SetupKeyStillRunsWhenPeerNeedsLogin(t *testing.T) { + s, _, _, username, _ := setupServerWithProfile(t) + s.rootCtx = internal.CtxInitState(context.Background()) + + s.isLoginRequiredFn = func(context.Context) (bool, error) { + return true, nil + } + + var keysTried []string + s.loginAttemptFn = func(_ context.Context, setupKey, _ string) (internal.StatusType, error) { + keysTried = append(keysTried, setupKey) + return "", nil + } + + setupKey := "A2C8E32F-AEB2-4B45-8FD3-8A0C1B2D3E4F" + resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username, SetupKey: setupKey}) + require.NoError(t, err, "the probe's outcome leaked out as the login result") + require.NotNil(t, resp) + require.Equal(t, []string{setupKey}, keysTried, "the setup key never reached the login attempt") + require.Nil(t, s.oauthAuthFlow.flow, "a setup-key login started an SSO flow") + + status, err := internal.CtxGetState(s.rootCtx).Status() + require.NoError(t, err) + require.Equal(t, internal.StatusIdle, status) +} + // breakProfilePrivateKey replaces the profile's private key with an unparseable // one, which makes any attempt to build a Management client fail on the spot. func breakProfilePrivateKey(t *testing.T, cfgPath string) { diff --git a/client/server/server.go b/client/server/server.go index 892c9c5de..eb6a8f2bc 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -140,6 +140,8 @@ type Server struct { // it to drive the login outcomes that need a server on the other end; // production leaves it nil, and every login goes through loginAttempt. loginAttemptFn func(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error) + + isLoginRequiredFn func(ctx context.Context) (bool, error) } type oauthAuthFlow struct { @@ -384,6 +386,21 @@ func (s *Server) attemptLogin(ctx context.Context, setupKey, jwtToken string) (i return s.loginAttempt(ctx, setupKey, jwtToken) } +func (s *Server) isLoginRequired(ctx context.Context) (bool, error) { + if s.isLoginRequiredFn != nil { + return s.isLoginRequiredFn(ctx) + } + + authClient, err := auth.NewAuth(ctx, s.config.PrivateKey, s.config.ManagementURL, s.config) + if err != nil { + log.Errorf("failed to create auth client: %v", err) + return false, err + } + defer authClient.Close() + + return authClient.IsLoginRequired(ctx) +} + // loginAttempt attempts to login using the provided information. It returns // StatusNeedsLogin when Management refused the peer's credentials and // StatusLoginFailed for every other failure, so callers can tell an @@ -640,22 +657,22 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro s.config = config s.mutex.Unlock() - loginStatus, err := s.attemptLogin(ctx, "", "") - if err == nil { - state.Set(internal.StatusIdle) - return &proto.LoginResponse{}, nil - } - - // Only an authentication refusal means the peer has to (re-)authenticate. - // Any other failure leaves the login undecided: Management unreachable, a + // A probe that errors leaves the login undecided: Management unreachable, a // restart mid-request, an internal error. Those are returned for the caller // to retry, because turning them into an SSO prompt asks the user to solve // something that is not theirs to solve, and a browser login cannot succeed - // while Management is unreachable anyway. - if loginStatus != internal.StatusNeedsLogin { - state.Set(loginStatus) + // while Management is unreachable anyway. Only Management refusing the + // peer's key is a decision, and IsLoginRequired reports that as + // needsLogin=true rather than an error. + needsLogin, err := s.isLoginRequired(ctx) + if err != nil { + state.Set(internal.StatusLoginFailed) return nil, err } + if !needsLogin { + state.Set(internal.StatusIdle) + return &proto.LoginResponse{}, nil + } if msg.SetupKey == "" { hint := "" From 564595d283197e0aeca6d08ac5e3b547ef41a14d Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 4 Aug 2026 14:02:09 +0000 Subject: [PATCH 30/34] [client, android] Fix profile account path test on Windows (#7057) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes The expected paths were hardcoded with Unix separators while profileAccountPathFor builds the result with filepath.Join, so the comparison failed on Windows. Derive the expectations with filepath.FromSlash to keep the test platform-independent. ## 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) - [ ] 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). ## 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 * **Tests** * Updated profile account path test expectations to use platform-appropriate path separators, improving test reliability across operating systems. --- client/android/profile_state_test.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/client/android/profile_state_test.go b/client/android/profile_state_test.go index 412435bc3..623e16c3b 100644 --- a/client/android/profile_state_test.go +++ b/client/android/profile_state_test.go @@ -16,17 +16,17 @@ func TestProfileAccountPathFor(t *testing.T) { { name: "default profile", configPath: "/data/data/io.netbird.client/files/netbird.cfg", - want: "/data/data/io.netbird.client/files/netbird.account.json", + want: filepath.FromSlash("/data/data/io.netbird.client/files/netbird.account.json"), }, { name: "id profile", configPath: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.json", - want: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.account.json", + want: filepath.FromSlash("/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.account.json"), }, { name: "legacy name-keyed profile is handled the same way", configPath: "/data/data/io.netbird.client/files/profiles/work.json", - want: "/data/data/io.netbird.client/files/profiles/work.account.json", + want: filepath.FromSlash("/data/data/io.netbird.client/files/profiles/work.account.json"), }, { name: "empty path is rejected", From f2d13b884ae835ef53a313823744dec1eed83f77 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 4 Aug 2026 14:02:31 +0000 Subject: [PATCH 31/34] [client] Fix session expired relogin (#7055) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes After the SSO session expires, the daemon tears the engine down permanently (management returns `PermissionDenied` → `runCancel()` → the retry loop exits for good). The "Session expired" dialog's Login button still drove the extend-session flow, which requires a live engine: the user completed the full browser SSO + 2FA round trip only to get `Failed to extend the session — engine is not initialised`, with no way out other than quitting and relaunching the client. Reproduce: 1. Log in on a desktop client with session expiration enabled (e.g. 16h TTL). 2. Let the session expire (e.g. leave the machine asleep overnight). 3. Wake it, click **Login** on the "Session expired" dialog, complete SSO + 2FA. 4. The error dialog appears and every retry fails the same way. Changes: - The expired branch of the session-expiration dialog now emits `trigger-login`, driving the full `Login → SSO → Up` sequence that rebuilds the client, instead of the extend flow (an expired session can no longer be extended). - `RequestExtendAuthSession` fails fast when the engine is already gone, so the browser/2FA round trip is not wasted on a doomed extend. - The expired tray row navigated the main window to `/#/login`, a route that does not exist and fell through to the main page without starting a login; it now emits `trigger-login` as well. ## 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) - [ ] 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). ## 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 - **Bug Fixes** - Improved session extension handling when the client engine is unavailable by prompting users to log in again. - Updated expired-session behavior to trigger the standard login flow, providing a more consistent sign-in experience. --- client/server/server.go | 3 +++ .../session/SessionExpirationDialog.tsx | 19 +++++++++++++++++-- client/ui/tray_session.go | 9 ++------- 3 files changed, 22 insertions(+), 9 deletions(-) diff --git a/client/server/server.go b/client/server/server.go index eb6a8f2bc..01778b8e0 100644 --- a/client/server/server.go +++ b/client/server/server.go @@ -1815,6 +1815,9 @@ func (s *Server) RequestExtendAuthSession( if connectClient == nil { return nil, gstatus.Errorf(codes.FailedPrecondition, "client is not running") } + if connectClient.Engine() == nil { + return nil, gstatus.Errorf(codes.FailedPrecondition, "session can no longer be extended, log in again to reconnect") + } hint := "" if msg.Hint != nil { diff --git a/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx b/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx index 10e71babb..ef8d6862f 100644 --- a/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx +++ b/client/ui/frontend/src/modules/session/SessionExpirationDialog.tsx @@ -11,7 +11,7 @@ import { DialogHeading } from "@/components/dialog/DialogHeading"; import { SquareIcon } from "@/components/SquareIcon"; import { Connection, Profiles as ProfilesSvc, Session, WindowManager } from "@bindings/services"; import { useAutoSizeWindow } from "@/hooks/useAutoSizeWindow"; -import { EVENT_BROWSER_LOGIN_CANCEL } from "@/lib/connection"; +import { EVENT_BROWSER_LOGIN_CANCEL, EVENT_TRIGGER_LOGIN } from "@/lib/connection"; import { errorDialog, formatErrorMessage } from "@/lib/errors.ts"; import { formatRemaining } from "@/lib/formatters"; @@ -131,6 +131,21 @@ export default function SessionExpirationDialog() { } }, [busy, t]); + const authenticate = useCallback(async () => { + if (busy) return; + setBusy(true); + try { + await Events.Emit(EVENT_TRIGGER_LOGIN); + await WindowManager.CloseSessionExpiration(); + } catch (e) { + setBusy(false); + await errorDialog({ + Title: t("connect.error.loginTitle"), + Message: formatErrorMessage(e), + }); + } + }, [busy, t]); + const logout = useCallback(async () => { if (busy) return; setBusy(true); @@ -185,7 +200,7 @@ export default function SessionExpirationDialog() { variant={"primary"} size={"md"} className={"w-full"} - onClick={stay} + onClick={expired ? authenticate : stay} disabled={busy} > {expired ? t("sessionExpiration.authenticate") : t("sessionExpiration.stay")} diff --git a/client/ui/tray_session.go b/client/ui/tray_session.go index 885fdb348..f25419894 100644 --- a/client/ui/tray_session.go +++ b/client/ui/tray_session.go @@ -27,11 +27,10 @@ const ( finalWarningCountdownSeconds = 120 ) -// handleSessionExpired notifies and brings the window forward so the frontend's /login route drives renewal. +// handleSessionExpired notifies and brings the window forward so the user can reconnect. func (t *Tray) handleSessionExpired() { t.notify(t.loc.T("notify.sessionExpired.title"), t.loc.T("notify.sessionExpired.body"), notifyIDSessionExpired) if t.window != nil { - t.window.SetURL("/#/login") t.window.Show() t.window.Focus() } @@ -308,11 +307,7 @@ func (t *Tray) openSessionExtendFlow() { } seconds := int(time.Until(deadline).Seconds()) if seconds <= 0 { - if t.window != nil { - t.window.SetURL("/#/login") - t.window.Show() - t.window.Focus() - } + t.app.Event.Emit(services.EventTriggerLogin) return } if t.svc.WindowManager == nil { From 2a61eac0474f75b8f48ed29f7fcbc85700aec705 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Tue, 4 Aug 2026 15:36:04 +0000 Subject: [PATCH 32/34] [client] Fix Linux tray right-click opening the main window (#7039) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Describe your changes On SNI hosts that report icon clicks via Activate (KDE Plasma, Waybar), Wails fired the left-click handler on every dbusmenu 'opened' event, so a right click opened the tray menu and immediately raised the main window, which stole focus and closed the menu. Host Left click Right click KDE Plasma, Waybar main window (Activate) menu (host-rendered) GNOME Shell + AppIndicator menu only menu only Minimal WMs via XEmbed host main window (Activate) XEmbed GTK popup Point the wails/v3 replace at the netbirdio fork (v3.0.0-beta.3 plus the fix): once the host has sent Activate/SecondaryActivate, a menu open no longer fires the click handler, while AppIndicator-only hosts that signal clicks solely via 'opened' keep the old 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) - [ ] 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). ## 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 ## Chores * Updated underlying application components to support compatibility and ongoing maintenance. * Clarified Linux tray interaction documentation, including platform-specific left- and right-click behavior and menu activation. * No user-facing features, workflow changes, or visual updates are included. * Existing Linux tray behavior remains unchanged. --- client/ui/tray_click_linux.go | 31 ++++++++++++++++++++----------- go.mod | 2 ++ go.sum | 4 ++-- 3 files changed, 24 insertions(+), 13 deletions(-) diff --git a/client/ui/tray_click_linux.go b/client/ui/tray_click_linux.go index 34a364fd9..95f5dfe85 100644 --- a/client/ui/tray_click_linux.go +++ b/client/ui/tray_click_linux.go @@ -4,17 +4,26 @@ package main // bindTrayClick wires the tray icon's left-click handler on Linux. // -// Both Linux click paths converge on Wails' linuxSystemTray.Activate, which -// fires the registered clickHandler: -// - Real SNI hosts (KDE Plasma, Waybar, GNOME Shell + AppIndicator) invoke -// org.kde.StatusNotifierItem.Activate over D-Bus on left-click. -// - The in-process StatusNotifierWatcher + XEmbed host used on minimal WMs -// (Fluxbox, i3, dwm, OpenBox) maps a Button1 press to that same Activate -// call itself (xembed_host_linux.go), so it routes through the same hook. -// Registering OnClick here therefore covers both paths with one handler — no -// changes to the watcher or XEmbed C code are needed. Left-click now opens the -// main window; right-click still opens the menu via Wails' default -// SecondaryActivate→OpenMenu handler (and the XEmbed GTK popup on minimal WMs). +// Expected behaviour per tray host: +// +// Host Left click Right click +// KDE Plasma, Waybar main window (Activate) menu (host-rendered) +// GNOME Shell + AppIndicator menu only menu only +// Minimal WMs via XEmbed host main window (Activate) XEmbed GTK popup +// +// OnClick fires only on org.kde.StatusNotifierItem.Activate — a real left +// click. KDE/Waybar send it over D-Bus; the in-process XEmbed host +// (xembed_host_linux.go) maps a Button1 press to the same Activate call. +// +// GNOME Shell + AppIndicator never sends Activate: it renders the dbusmenu +// on ANY click and only reports the menu opening via dbusmenu +// Event("opened"). Upstream Wails treated that event as a click, so on GNOME +// both buttons raised the main window on top of the menu, and on KDE/Waybar +// a right click raised it over the freshly opened menu. The netbirdio/wails +// fork (go.mod replace) drops that heuristic: a menu open never fires +// OnClick. On GNOME the main window is reached via the "Open NetBird" menu +// entry; left-click-opens-window is not achievable there anyway, since the +// host always opens the menu itself. // // We do NOT register OnDoubleClick: Wails' Linux SNI backend never fires it // (unlike Windows). And we deliberately skip AttachWindow — it plus Wails3's diff --git a/go.mod b/go.mod index 44d9e908f..94b2eb059 100644 --- a/go.mod +++ b/go.mod @@ -339,3 +339,5 @@ replace github.com/dexidp/dex => github.com/netbirdio/dex v0.244.1-0.20260716205 replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1 replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0 + +replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4 diff --git a/go.sum b/go.sum index 762038e9e..91ce14073 100644 --- a/go.sum +++ b/go.sum @@ -490,6 +490,8 @@ github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9ax github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ= github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4 h1:UKztc3QjWvzU5DZk+uYaOWN0x62NSe/pkxuPvzqZIy4= +github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4/go.mod h1:BzATbK71VFikMMMCo434wAi0QcaI03P+xeaWgDvQvjw= github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a h1:3CWK+yTvRKOcC0Q8VCTGy4l60TEb27CQVS7LkMxwjmw= github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw= github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A= @@ -660,8 +662,6 @@ github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IU github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok= github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g= github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds= -github.com/wailsapp/wails/v3 v3.0.0-beta.3 h1:BrcZunEBVucncRx+xgkk9TzlXU4qc0ygJuEhKAAGaeA= -github.com/wailsapp/wails/v3 v3.0.0-beta.3/go.mod h1:BzATbK71VFikMMMCo434wAi0QcaI03P+xeaWgDvQvjw= github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU= github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA= github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM= From 2afa69b622d4cdd1eba306eef4679eb8c2af1d9d Mon Sep 17 00:00:00 2001 From: Misha Bragin Date: Tue, 4 Aug 2026 18:04:22 +0200 Subject: [PATCH 33/34] [management] prevent dangling group refs in agent-network ACLs. (#7060) Block deleting a group referenced as a source group by an agent network policy, and drop unresolvable groups from synthesised private-service ACLs. A deleted group survived in agent_network_policies.source_groups and was carried into the injected in-memory policy, where network-map assembly resolved it to a nil group and panicked on every proxy peer sync. --- management/server/group.go | 25 ++++++++++++++++++++++++ management/server/group_test.go | 31 ++++++++++++++++++++++++++++++ management/server/types/account.go | 26 ++++++++++++++++++++++--- 3 files changed, 79 insertions(+), 3 deletions(-) diff --git a/management/server/group.go b/management/server/group.go index dab891f2a..e6f748d02 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -6,6 +6,7 @@ import ( "fmt" "slices" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" "github.com/rs/xid" log "github.com/sirupsen/logrus" @@ -744,6 +745,10 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty return &GroupLinkError{"network router", linkedRouter.ID} } + if isLinked, linkedPolicy := isGroupLinkedToAgentNetworkPolicy(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"agent network policy", linkedPolicy.Name} + } + return checkGroupLinkedToSettings(ctx, transaction, group) } @@ -875,6 +880,26 @@ func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store, return false, nil } +// isGroupLinkedToAgentNetworkPolicy checks if a group is used as a source group by any +// agent network policy in the account. +func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.Policy) { + policies, err := transaction.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving agent network policies while checking group linkage: %v", err) + return false, nil + } + + for _, policy := range policies { + if policy == nil { + continue + } + if slices.Contains(policy.SourceGroups, groupID) { + return true, policy + } + } + return false, nil +} + // areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers. // It fetches each collection once and checks all groupIDs against them in memory. func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) { diff --git a/management/server/group_test.go b/management/server/group_test.go index 22fda2671..17bb59d6a 100644 --- a/management/server/group_test.go +++ b/management/server/group_test.go @@ -18,6 +18,7 @@ import ( "golang.org/x/exp/maps" nbdns "github.com/netbirdio/netbird/dns" + agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/networks" "github.com/netbirdio/netbird/management/server/networks/resources" @@ -125,6 +126,11 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) { "grp-for-integration", "only service users with admin power can delete integration group", }, + { + "agent network policy", + "grp-for-agent-network-policy", + "agent network policy", + }, } for _, testCase := range testCases { @@ -218,6 +224,11 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) { groupIDs: []string{"grp-for-integration"}, expectedReasons: []string{"only service users with admin power can delete integration group"}, }, + { + name: "agent network policy", + groupIDs: []string{"grp-for-agent-network-policy"}, + expectedReasons: []string{"agent network policy"}, + }, { name: "successfully delete multiple groups", groupIDs: []string{"group-1", "group-2"}, @@ -406,6 +417,14 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t Peers: make([]string, 0), } + groupForAgentNetworkPolicy := &types.Group{ + ID: "grp-for-agent-network-policy", + AccountID: "account-id", + Name: "Group for agent network policies", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + routeResource := &route.Route{ ID: "example route", Groups: []string{groupForRoute.ID}, @@ -461,6 +480,18 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForSetupKeys) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForUsers) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForIntegration) + _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkPolicy) + + agentNetworkPolicy := &agentNetworkTypes.Policy{ + ID: "example agent network policy", + AccountID: accountID, + Name: "Example agent network policy", + Enabled: true, + SourceGroups: []string{groupForAgentNetworkPolicy.ID}, + } + if err := am.Store.SaveAgentNetworkPolicy(context.Background(), agentNetworkPolicy); err != nil { + return nil, nil, err + } acc, err := am.Store.GetAccount(context.Background(), account.Id) if err != nil { diff --git a/management/server/types/account.go b/management/server/types/account.go index 1a3a30544..474825281 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -1707,14 +1707,34 @@ func (a *Account) injectPrivateServicePolicies(svc *service.Service, proxyPeers if len(proxyPeers) == 0 { return } + // A service's AccessGroups can name groups that no longer exist — persisted + // services and the agent-network synthesiser both carry the ids verbatim from + // their own state. An unresolvable source authorises nothing, so drop it here + // rather than let the network-map assembly resolve it to a nil group. + sources := a.existingGroupIDs(svc.AccessGroups) + if len(sources) == 0 { + return + } for _, proxyPeer := range proxyPeers { - a.Policies = append(a.Policies, a.createPrivateServicePolicy(svc, proxyPeer)) + a.Policies = append(a.Policies, a.createPrivateServicePolicy(svc, proxyPeer, sources)) } } -func (a *Account) createPrivateServicePolicy(svc *service.Service, proxyPeer *nbpeer.Peer) *Policy { +// existingGroupIDs returns the subset of groupIDs that resolve to a group in the account, +// preserving the input order. +func (a *Account) existingGroupIDs(groupIDs []string) []string { + out := make([]string, 0, len(groupIDs)) + for _, groupID := range groupIDs { + if _, ok := a.Groups[groupID]; ok { + out = append(out, groupID) + } + } + return out +} + +func (a *Account) createPrivateServicePolicy(svc *service.Service, proxyPeer *nbpeer.Peer, accessGroups []string) *Policy { policyID := fmt.Sprintf("private-access-%s-%s", svc.ID, proxyPeer.ID) - sources := append([]string(nil), svc.AccessGroups...) + sources := append([]string(nil), accessGroups...) return &Policy{ ID: policyID, Name: fmt.Sprintf("Private Access to %s", svc.Name), From 6526fc2bec09174cb713a0220f430109c169dba1 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Wed, 5 Aug 2026 03:24:20 +0900 Subject: [PATCH 34/34] [management] Prevent deleting groups referenced by reverse proxy services (#7062) ## Describe your changes A group could be deleted while a reverse proxy service still referenced it, silently breaking the service's access control: private services list groups in `access_groups` as the peer allowlist, and SSO bearer auth distributes tokens to `distribution_groups`. Group deletion now runs through the same linkage validation as routes, policies, and agent network policies: deleting a group that backs a private service allowlist or an enabled bearer-auth distribution list fails with a `GroupLinkError` naming the service domain. Disabled bearer configs and stale `access_groups` on non-private services are inert and do not block deletion. Tests cover both linked cases in single and bulk deletion, and pin the non-blocking cases. The test account seeds decoy services ahead of the linked ones so the check is proven to scan the full service list. --- management/server/group.go | 25 ++++++ management/server/group_test.go | 140 ++++++++++++++++++++++++++++++++ 2 files changed, 165 insertions(+) diff --git a/management/server/group.go b/management/server/group.go index e6f748d02..33870f25e 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -7,6 +7,7 @@ import ( "slices" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/rs/xid" log "github.com/sirupsen/logrus" @@ -745,6 +746,10 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty return &GroupLinkError{"network router", linkedRouter.ID} } + if isLinked, linkedService := isGroupLinkedToReverseProxyService(ctx, transaction, group.AccountID, group.ID); isLinked { + return &GroupLinkError{"reverse proxy service", linkedService.Domain} + } + if isLinked, linkedPolicy := isGroupLinkedToAgentNetworkPolicy(ctx, transaction, group.AccountID, group.ID); isLinked { return &GroupLinkError{"agent network policy", linkedPolicy.Name} } @@ -880,6 +885,26 @@ func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store, return false, nil } +// isGroupLinkedToReverseProxyService checks if a group is used as an access group +// of a private reverse proxy service or as a bearer-auth distribution group. +func isGroupLinkedToReverseProxyService(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *service.Service) { + services, err := transaction.GetAccountServices(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving reverse proxy services while checking group linkage: %v", err) + return false, nil + } + + for _, svc := range services { + if svc.Private && slices.Contains(svc.AccessGroups, groupID) { + return true, svc + } + if svc.Auth.BearerAuth != nil && svc.Auth.BearerAuth.Enabled && slices.Contains(svc.Auth.BearerAuth.DistributionGroups, groupID) { + return true, svc + } + } + return false, nil +} + // isGroupLinkedToAgentNetworkPolicy checks if a group is used as a source group by any // agent network policy in the account. func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.Policy) { diff --git a/management/server/group_test.go b/management/server/group_test.go index 17bb59d6a..deeec61d5 100644 --- a/management/server/group_test.go +++ b/management/server/group_test.go @@ -19,6 +19,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/networks" "github.com/netbirdio/netbird/management/server/networks/resources" @@ -131,6 +132,16 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) { "grp-for-agent-network-policy", "agent network policy", }, + { + "reverse proxy private service access group", + "grp-for-rp-private", + "reverse proxy service", + }, + { + "reverse proxy bearer distribution group", + "grp-for-rp-bearer", + "reverse proxy service", + }, } for _, testCase := range testCases { @@ -229,6 +240,12 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) { groupIDs: []string{"grp-for-agent-network-policy"}, expectedReasons: []string{"agent network policy"}, }, + { + name: "reverse proxy services", + groupIDs: []string{"grp-for-rp-private", "grp-for-rp-bearer"}, + expectedReasons: []string{"reverse proxy service", "reverse proxy service"}, + expectedNotDeleted: []string{"grp-for-rp-private", "grp-for-rp-bearer"}, + }, { name: "successfully delete multiple groups", groupIDs: []string{"group-1", "group-2"}, @@ -296,6 +313,65 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) { } } +func TestDefaultAccountManager_DeleteGroupUnlinkedFromReverseProxyService(t *testing.T) { + am, _, err := createManager(t) + require.NoError(t, err, "Failed to create account manager") + + _, account, err := initTestGroupAccount(am) + require.NoError(t, err, "Failed to init testing account") + + deletableGroups := []*types.Group{ + { + ID: "grp-rp-bearer-disabled", + AccountID: account.Id, + Name: "Group only in a disabled bearer auth", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + }, + { + ID: "grp-rp-nonprivate-access", + AccountID: account.Id, + Name: "Group only in a non-private service's access groups", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + }, + } + for _, group := range deletableGroups { + require.NoError(t, am.CreateGroup(context.Background(), account.Id, groupAdminUserID, group)) + } + + // Disabled bearer auth and stale access groups on a non-private service + // are inert configuration and must not block group deletion. + services := []*rpservice.Service{ + { + ID: "rp-svc-bearer-disabled", + AccountID: account.Id, + Domain: "bearer-disabled.services.example.com", + Auth: rpservice.AuthConfig{ + BearerAuth: &rpservice.BearerAuthConfig{ + Enabled: false, + DistributionGroups: []string{"grp-rp-bearer-disabled"}, + }, + }, + }, + { + ID: "rp-svc-nonprivate-access", + AccountID: account.Id, + Domain: "nonprivate.services.example.com", + Private: false, + AccessGroups: []string{"grp-rp-nonprivate-access"}, + }, + } + for _, svc := range services { + require.NoError(t, am.Store.CreateService(context.Background(), svc)) + } + + for _, group := range deletableGroups { + err = am.DeleteGroup(context.Background(), account.Id, groupAdminUserID, group.ID) + assert.NoError(t, err, "group %s is not referenced by an active reverse proxy gate and should be deletable", group.ID) + } +} + func TestDefaultAccountManager_DeleteGroupLinkedToFlowGroup(t *testing.T) { am, _, err := createManager(t) require.NoError(t, err) @@ -425,6 +501,22 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t Peers: make([]string, 0), } + groupForRPPrivate := &types.Group{ + ID: "grp-for-rp-private", + AccountID: "account-id", + Name: "Group for private reverse proxy service", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + + groupForRPBearer := &types.Group{ + ID: "grp-for-rp-bearer", + AccountID: "account-id", + Name: "Group for bearer reverse proxy service", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + routeResource := &route.Route{ ID: "example route", Groups: []string{groupForRoute.ID}, @@ -481,6 +573,8 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForUsers) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForIntegration) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkPolicy) + _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPPrivate) + _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPBearer) agentNetworkPolicy := &agentNetworkTypes.Policy{ ID: "example agent network policy", @@ -493,6 +587,52 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t return nil, nil, err } + // The decoy services are created first so the linkage check has to scan + // past services that do not reference the groups under test. + rpServices := []*rpservice.Service{ + { + ID: "rp-svc-private-decoy", + AccountID: accountID, + Domain: "private-decoy.services.example.com", + Private: true, + AccessGroups: []string{"unrelated-group"}, + }, + { + ID: "rp-svc-bearer-decoy", + AccountID: accountID, + Domain: "bearer-decoy.services.example.com", + Auth: rpservice.AuthConfig{ + BearerAuth: &rpservice.BearerAuthConfig{ + Enabled: true, + DistributionGroups: []string{"unrelated-group"}, + }, + }, + }, + { + ID: "rp-svc-private", + AccountID: accountID, + Domain: "private.services.example.com", + Private: true, + AccessGroups: []string{groupForRPPrivate.ID}, + }, + { + ID: "rp-svc-bearer", + AccountID: accountID, + Domain: "bearer.services.example.com", + Auth: rpservice.AuthConfig{ + BearerAuth: &rpservice.BearerAuthConfig{ + Enabled: true, + DistributionGroups: []string{groupForRPBearer.ID}, + }, + }, + }, + } + for _, svc := range rpServices { + if err := am.Store.CreateService(context.Background(), svc); err != nil { + return nil, nil, err + } + } + acc, err := am.Store.GetAccount(context.Background(), account.Id) if err != nil { return nil, nil, err