Compare commits

..

3 Commits

Author SHA1 Message Date
bcmmbaga
6725b02cbb Merge branch 'main' into add-atomic-cache-ops 2026-07-22 20:20:14 +03:00
bcmmbaga
da19dcf480 prevent concurrent JWT reuse 2026-07-22 20:08:48 +03:00
bcmmbaga
6426d6f03f add atomic SetNX cache operation
Split the memory and Redis cache implementations into separate files and
provide backend-native atomic set-if-absent support.
2026-07-22 20:03:03 +03:00
142 changed files with 1136 additions and 5574 deletions

View File

@@ -5,13 +5,6 @@ on:
schedule:
- cron: "0 3 * * *"
workflow_dispatch:
inputs:
bedrock_model:
description: >-
Bedrock inference-profile id to drive the matrix with, exactly as
AWS issues it. Leave empty for the Sonnet 4.6 default.
required: false
default: ""
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -58,9 +51,6 @@ jobs:
# token (and URL, for gateways) is unset, so partial coverage is fine.
OPENAI_TOKEN: ${{ secrets.E2E_OPENAI_TOKEN }}
ANTHROPIC_TOKEN: ${{ secrets.E2E_ANTHROPIC_TOKEN }}
# Moonshot AI platform key (platform.kimi.ai); drives both Kimi wire
# shapes (OpenAI /v1 and Anthropic /anthropic) through kimi_api.
KIMI_TOKEN: ${{ secrets.E2E_KIMI_TOKEN }}
VERCEL_URL: ${{ secrets.E2E_VERCEL_URL }}
VERCEL_TOKEN: ${{ secrets.E2E_VERCEL_TOKEN }}
OPENROUTER_URL: ${{ secrets.E2E_OPENROUTER_URL }}
@@ -69,8 +59,6 @@ jobs:
CLOUDFLARE_TOKEN: ${{ secrets.E2E_CLOUDFLARE_TOKEN }}
AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.E2E_AWS_BEARER_TOKEN_BEDROCK }}
AWS_REGION: ${{ secrets.E2E_AWS_REGION }}
# Bedrock model override: dispatch input wins, then the repo variable, else the test default.
AWS_BEDROCK_MODEL: ${{ inputs.bedrock_model || vars.E2E_AWS_BEDROCK_MODEL }}
# Vertex (Anthropic-on-Vertex): SA + project required; region defaults
# to "global", model to a pinned claude snapshot.
GOOGLE_VERTEX_SA_BASE64: ${{ secrets.E2E_GOOGLE_VERTEX_SA_BASE64 }}

View File

@@ -273,8 +273,8 @@ dockers_v2:
- netbirdio/netbird
- ghcr.io/netbirdio/netbird
tags:
- "{{ .Version }}-rootless"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}rootless-latest{{ end }}"
- "v{{ .Version }}-rootless"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
dockerfile: client/Dockerfile-rootless
extra_files:
- client/netbird-entrypoint.sh

View File

@@ -464,8 +464,6 @@ func Test_RemovePeer(t *testing.T) {
}
func Test_ConnectPeers(t *testing.T) {
t.Setenv("NB_DISABLE_EBPF_WG_PROXY", "true")
peer1ifaceName := fmt.Sprintf("utun%d", WgIntNumber+400)
peer1wgIP := netip.MustParsePrefix("10.99.99.17/30")
peer1Key, _ := wgtypes.GeneratePrivateKey()
@@ -507,8 +505,12 @@ func Test_ConnectPeers(t *testing.T) {
t.Fatal(err)
}
localIP1 := "127.0.0.1"
peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP1, peer1wgPort))
localIP, err := getLocalIP()
if err != nil {
t.Fatal(err)
}
peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer1wgPort))
if err != nil {
t.Fatal(err)
}
@@ -544,8 +546,7 @@ func Test_ConnectPeers(t *testing.T) {
t.Fatal(err)
}
localIP2 := "127.0.0.1"
peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP2, peer2wgPort))
peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer2wgPort))
if err != nil {
t.Fatal(err)
}
@@ -568,17 +569,17 @@ func Test_ConnectPeers(t *testing.T) {
if err != nil {
t.Fatal(err)
}
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop
// here starves the wireguard-go goroutines that process the handshake, so
// poll on a ticker instead and yield the CPU between checks. WireGuard also
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which
// is why the overall wait can occasionally stretch to tens of seconds.
// todo: investigate why in some tests execution we need 30s
timeout := 30 * time.Second
timeoutChannel := time.After(timeout)
ticker := time.NewTicker(500 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-timeoutChannel:
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
default:
}
peer, gpErr := getPeer(peer1ifaceName, peer2Key.PublicKey().String())
if gpErr != nil {
t.Fatal(gpErr)
@@ -587,12 +588,6 @@ func Test_ConnectPeers(t *testing.T) {
t.Log("peers successfully handshake")
break
}
select {
case <-timeoutChannel:
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
case <-ticker.C:
}
}
}
@@ -620,3 +615,28 @@ func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
}
return wgtypes.Peer{}, fmt.Errorf("peer not found")
}
func getLocalIP() (string, error) {
// Get all interfaces
addrs, err := net.InterfaceAddrs()
if err != nil {
return "", err
}
for _, addr := range addrs {
ipNet, ok := addr.(*net.IPNet)
if !ok {
continue
}
if ipNet.IP.IsLoopback() {
continue
}
if ipNet.IP.To4() == nil {
continue
}
return ipNet.IP.String(), nil
}
return "", fmt.Errorf("no local IP found")
}

View File

@@ -49,21 +49,11 @@ type ConnMgr struct {
// engine.syncMsgMux; all other reads stay under engine.syncMsgMux only.
lazyConnMgrMu sync.RWMutex
// reconcileRoutedIPs re-applies a peer's routed allowed IPs after its lazy wake endpoint is
// (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile.
reconcileRoutedIPs func(peerKey string) error
wg sync.WaitGroup
lazyCtx context.Context
lazyCtxCancel context.CancelFunc
}
// SetRoutedIPsReconciler injects the callback used to re-apply a peer's routed allowed IPs when
// its lazy wake endpoint is (re)armed. Must be called before the lazy manager starts.
func (e *ConnMgr) SetRoutedIPsReconciler(fn func(peerKey string) error) {
e.reconcileRoutedIPs = fn
}
func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr {
e := &ConnMgr{
peerStore: peerStore,
@@ -301,7 +291,6 @@ func (e *ConnMgr) Close() {
func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
cfg := manager.Config{
InactivityThreshold: inactivityThresholdEnv(),
ReconcileAllowedIPs: e.reconcileRoutedIPs,
}
e.lazyConnMgrMu.Lock()

View File

@@ -52,14 +52,11 @@ int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) {
if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) {
udp->dest = dns_port;
// Clear the now-stale checksum; zero means "not computed" for IPv4.
udp->check = 0;
return XDP_PASS;
}
if (udp->source == dns_port && ip->saddr == dns_ip) {
udp->source = GENERAL_DNS_PORT;
udp->check = 0;
return XDP_PASS;
}

View File

@@ -50,11 +50,5 @@ int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) {
__be16 new_dst_port = htons(proxy_port);
udp->dest = new_dst_port;
udp->source = new_src_port;
// The ports are covered by the UDP checksum. This is an IPv4 loopback hop
// and the payload is already integrity-protected, so clear the checksum (a
// zero UDP checksum means "not computed" for IPv4) rather than leave a
// stale value the kernel would drop as UDP_CSUM.
udp->check = 0;
return XDP_PASS;
}

View File

@@ -663,12 +663,6 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
iceCfg := e.createICEConfig()
e.connMgr = NewConnMgr(e.config, e.statusRecorder, e.peerStore, wgIface)
e.connMgr.SetRoutedIPsReconciler(func(peerKey string) error {
if e.routeManager == nil {
return nil
}
return e.routeManager.ReconcilePeerAllowedIPs(peerKey)
})
e.connMgr.Start(e.ctx)
// Wire DNS-time lazy-connection warm-up now that the connection manager

View File

@@ -29,11 +29,6 @@ type managedPeer struct {
type Config struct {
InactivityThreshold *time.Duration
// ReconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is
// armed. The activity listener creates the wake peer with the overlay /32 only; without the
// routed prefixes WireGuard would not steer subnet-bound traffic to the wake endpoint, so an
// idle routing peer could never be woken by that traffic. Optional; nil disables the reconcile.
ReconcileAllowedIPs func(peerKey string) error
}
// Manager manages lazy connections
@@ -61,9 +56,6 @@ type Manager struct {
peerToHAGroups map[string][]route.HAUniqueID // peer ID -> HA groups they belong to
haGroupToPeers map[route.HAUniqueID][]string // HA group -> peer IDs in the group
routesMu sync.RWMutex
// reconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is armed.
reconcileAllowedIPs func(peerKey string) error
}
// NewManager creates a new lazy connection manager
@@ -81,7 +73,6 @@ func NewManager(config Config, engineCtx context.Context, peerStore *peerstore.S
activityManager: activity.NewManager(wgIface),
peerToHAGroups: make(map[string][]route.HAUniqueID),
haGroupToPeers: make(map[route.HAUniqueID][]string),
reconcileAllowedIPs: config.ReconcileAllowedIPs,
}
if wgIface.IsUserspaceBind() {
@@ -210,7 +201,7 @@ func (m *Manager) AddPeer(peerCfg lazyconn.PeerConfig) (bool, error) {
return false, nil
}
if err := m.armActivityListener(peerCfg); err != nil {
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
return false, err
}
@@ -297,7 +288,7 @@ func (m *Manager) DeactivatePeer(peerID peerid.ConnID) {
m.inactivityManager.RemovePeer(mp.peerCfg.PublicKey)
if err := m.armActivityListener(*mp.peerCfg); err != nil {
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
return
}
@@ -474,31 +465,6 @@ func (m *Manager) close() {
}
// shouldDeferIdleForHA checks if peer should stay connected due to HA group requirements
// armRoutedAllowedIPs re-applies the peer's routed allowed IPs onto its freshly armed wake
// endpoint. The activity listener creates the wake peer with the overlay /32 only, so without
// this the routed prefixes would be missing and traffic to a routed subnet could not wake the
// idle routing peer. It is a no-op when no reconciler is configured.
// armActivityListener (re)arms the peer's wake endpoint via the activity manager and then
// re-applies its routed allowed IPs, so traffic to a routed subnet can wake an idle routing
// peer. The routed prefixes must be re-applied after the wake endpoint exists because the
// listener creates it with the overlay /32 only.
func (m *Manager) armActivityListener(peerCfg lazyconn.PeerConfig) error {
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
return err
}
m.armRoutedAllowedIPs(&peerCfg)
return nil
}
func (m *Manager) armRoutedAllowedIPs(peerCfg *lazyconn.PeerConfig) {
if m.reconcileAllowedIPs == nil {
return
}
if err := m.reconcileAllowedIPs(peerCfg.PublicKey); err != nil {
peerCfg.Log.Errorf("failed to reconcile routed allowed IPs on wake endpoint: %v", err)
}
}
func (m *Manager) shouldDeferIdleForHA(inactivePeers map[string]struct{}, peerID string) bool {
m.routesMu.RLock()
defer m.routesMu.RUnlock()
@@ -611,7 +577,7 @@ func (m *Manager) onPeerInactivityTimedOut(peerIDs map[string]struct{}) {
mp.peerCfg.Log.Infof("start activity monitor")
if err := m.armActivityListener(*mp.peerCfg); err != nil {
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
continue
}

View File

@@ -95,7 +95,7 @@ func (d *DnsInterceptor) RemoveRoute() error {
// AllowedIPs should use real IPs
if d.currentPeerKey != "" {
if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil {
if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err))
}
}
@@ -172,7 +172,7 @@ func (d *DnsInterceptor) removeAllowedIP(realPrefix netip.Prefix) error {
}
// AllowedIPs use real IPs
if _, err := d.allowedIPsRefcounter.Decrement(realPrefix, d.currentPeerKey); err != nil {
if _, err := d.allowedIPsRefcounter.Decrement(realPrefix); err != nil {
return fmt.Errorf("remove allowed IP %s: %v", realPrefix, err)
}
@@ -205,7 +205,7 @@ func (d *DnsInterceptor) RemoveAllowedIPs() error {
for _, prefixes := range d.interceptedDomains {
for _, prefix := range prefixes {
// AllowedIPs use real IPs
if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil {
if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err))
}
}

View File

@@ -135,7 +135,7 @@ func (r *Route) RemoveAllowedIPs() error {
var merr *multierror.Error
for _, domainPrefixes := range r.dynamicDomains {
for _, prefix := range domainPrefixes {
if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil {
if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err))
}
}
@@ -320,7 +320,7 @@ func (r *Route) removeRoutes(prefixes []netip.Prefix) ([]netip.Prefix, error) {
merr = multierror.Append(merr, fmt.Errorf("remove dynamic route for IP %s: %w", prefix, err))
}
if r.currentPeerKey != "" {
if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil {
if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err))
}
}

View File

@@ -61,7 +61,6 @@ type Manager interface {
InitialRouteRange() []string
SetFirewall(firewall.Manager) error
SetDNSForwarderPort(port uint16)
ReconcilePeerAllowedIPs(peerKey string) error
Stop(stateManager *statemanager.Manager)
}
@@ -216,7 +215,7 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) {
)
}
m.allowedIPsRefCounter = refcounter.NewAllowedIPs(
m.allowedIPsRefCounter = refcounter.New(
func(prefix netip.Prefix, peerKey string) (string, error) {
// save peerKey to use it in the remove function
return peerKey, m.wgInterface.AddAllowedIP(peerKey, prefix)
@@ -233,30 +232,6 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) {
)
}
// ReconcilePeerAllowedIPs re-applies every routed allowed IP currently tracked for the peer
// onto the WireGuard device. The allowed-IP refcounter only calls its AddFunc (which pushes to
// the device) on a prefix's 0->1 transition, so a peer whose device entry was rebuilt without a
// matching refcounter change — e.g. a lazy connection cycling through idle->wake, which recreates
// the WireGuard peer with the overlay /32 only — ends up missing routed prefixes the refcounter
// still considers installed, and nothing retries. Calling this when the peer's WireGuard entry is
// (re)created restores convergence. It is add-only and idempotent: AddAllowedIP is update-only, so
// prefixes are re-added to an existing peer and an absent peer is left untouched.
func (m *DefaultManager) ReconcilePeerAllowedIPs(peerKey string) error {
if m.allowedIPsRefCounter == nil {
return nil
}
return m.allowedIPsRefCounter.ReapplyMatching(
func(out string) bool { return out == peerKey },
func(prefix netip.Prefix) error {
if err := m.wgInterface.AddAllowedIP(peerKey, prefix); err != nil {
return fmt.Errorf("add allowed IP %s for peer %s: %w", prefix, peerKey, err)
}
return nil
},
)
}
// Init sets up the routing
func (m *DefaultManager) Init() error {
m.routeSelector = m.initSelector()

View File

@@ -112,11 +112,6 @@ func (m *MockManager) SetFirewall(firewall.Manager) error {
func (m *MockManager) SetDNSForwarderPort(port uint16) {
}
// ReconcilePeerAllowedIPs mock implementation of ReconcilePeerAllowedIPs from Manager interface
func (m *MockManager) ReconcilePeerAllowedIPs(peerKey string) error {
return nil
}
// Stop mock implementation of Stop from Manager interface
func (m *MockManager) Stop(stateManager *statemanager.Manager) {
if m.StopFunc != nil {

View File

@@ -1,90 +0,0 @@
//go:build !windows
package routemanager
import (
"net"
"net/netip"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/tun/netstack"
"github.com/netbirdio/netbird/client/iface/device"
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/routemanager/refcounter"
)
// reconcileWGMock is a minimal iface.WGIface that only records AddAllowedIP calls; every other
// method is an inert stub because ReconcilePeerAllowedIPs exercises none of them.
type reconcileWGMock struct {
mu sync.Mutex
adds map[string][]netip.Prefix
}
func (m *reconcileWGMock) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
m.mu.Lock()
defer m.mu.Unlock()
if m.adds == nil {
m.adds = map[string][]netip.Prefix{}
}
m.adds[peerKey] = append(m.adds[peerKey], allowedIP)
return nil
}
func (m *reconcileWGMock) added(peerKey string) []netip.Prefix {
m.mu.Lock()
defer m.mu.Unlock()
return m.adds[peerKey]
}
func (m *reconcileWGMock) RemoveAllowedIP(string, netip.Prefix) error { return nil }
func (m *reconcileWGMock) Name() string { return "utun-test" }
func (m *reconcileWGMock) Address() wgaddr.Address { return wgaddr.Address{} }
func (m *reconcileWGMock) ToInterface() *net.Interface { return nil }
func (m *reconcileWGMock) IsUserspaceBind() bool { return false }
func (m *reconcileWGMock) GetFilter() device.PacketFilter { return nil }
func (m *reconcileWGMock) GetDevice() *device.FilteredDevice { return nil }
func (m *reconcileWGMock) GetNet() *netstack.Net { return nil }
// TestReconcilePeerAllowedIPs verifies the declarative reconcile re-applies every routed prefix
// tracked for the peer (self-heal, independent of refcount level) and stays scoped to that peer.
func TestReconcilePeerAllowedIPs(t *testing.T) {
wg := &reconcileWGMock{}
m := &DefaultManager{wgInterface: wg}
m.allowedIPsRefCounter = refcounter.NewAllowedIPs(
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
func(netip.Prefix, string) error { return nil },
)
peerA1 := netip.MustParsePrefix("10.0.0.0/24")
peerA2 := netip.MustParsePrefix("10.1.0.0/24")
peerB1 := netip.MustParsePrefix("10.2.0.0/24")
for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} {
_, err := m.allowedIPsRefCounter.Increment(prefix, peer)
require.NoError(t, err)
}
// Extra reference: reconcile must still re-apply the prefix even though its refcount never
// hit 0 again (the exact case the plain incremental path skips).
_, err := m.allowedIPsRefCounter.Increment(peerA1, "peerA")
require.NoError(t, err)
require.NoError(t, m.ReconcilePeerAllowedIPs("peerA"))
assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, wg.added("peerA"),
"reconcile must re-apply all routed prefixes of the peer")
assert.Empty(t, wg.added("peerB"), "reconcile must not touch another peer's prefixes")
}
// TestReconcilePeerAllowedIPsNoCounter verifies reconcile is a safe no-op before the refcounter is
// set up.
func TestReconcilePeerAllowedIPsNoCounter(t *testing.T) {
wg := &reconcileWGMock{}
m := &DefaultManager{wgInterface: wg}
require.NoError(t, m.ReconcilePeerAllowedIPs("peerA"))
assert.Empty(t, wg.added("peerA"))
}

View File

@@ -1,206 +0,0 @@
package refcounter
import (
"errors"
"fmt"
"net/netip"
"sort"
"sync"
"github.com/hashicorp/go-multierror"
nberrors "github.com/netbirdio/netbird/client/errors"
)
// allowedIPsEntry holds the per-peer reference counts for a single prefix and which peer is
// currently installed in WireGuard. WireGuard allows a prefix on exactly one peer, so at most
// one peer is active at a time even when several peers reference the prefix.
type allowedIPsEntry struct {
// peers maps a peerKey to the number of references holding the prefix for that peer.
peers map[string]int
// active is the peerKey currently installed in WireGuard for this prefix ("" if none).
active string
// total is the sum of all per-peer reference counts (kept in sync with peers).
total int
}
// AllowedIPsRefCounter is a peer-aware reference counter for WireGuard AllowedIPs.
//
// The generic Counter keys only by prefix and remembers a single Out value set by the first
// caller, which it never changes. That is wrong for AllowedIPs: two independent watchers (or
// multiple resolved domains) can reference the same prefix through different peers, and when the
// peer currently installed in WireGuard releases its last reference the prefix must be handed over
// to a surviving peer instead of being left pointing at the released one.
//
// It calls add/remove (which program WireGuard) only on the transitions that matter:
// - add on the first reference for a prefix, or when swapping the active peer;
// - remove on the last reference for a prefix, or on the old peer during a swap.
type AllowedIPsRefCounter struct {
mu sync.Mutex
entries map[netip.Prefix]*allowedIPsEntry
add AddFunc[netip.Prefix, string, string]
remove RemoveFunc[netip.Prefix, string]
}
// NewAllowedIPs creates a new peer-aware AllowedIPs reference counter.
// add programs a prefix on a peer in WireGuard and returns the peerKey to store as the active peer.
// remove unprograms the prefix from the given peer.
func NewAllowedIPs(add AddFunc[netip.Prefix, string, string], remove RemoveFunc[netip.Prefix, string]) *AllowedIPsRefCounter {
return &AllowedIPsRefCounter{
entries: map[netip.Prefix]*allowedIPsEntry{},
add: add,
remove: remove,
}
}
// Increment adds a reference to prefix for peerKey. WireGuard is programmed only for the first
// reference to a prefix; while a different peer is already installed the prefix is left with it
// (first peer wins, HA at the WireGuard layer is not possible) and only the reference count is kept.
func (rm *AllowedIPsRefCounter) Increment(prefix netip.Prefix, peerKey string) (Ref[string], error) {
rm.mu.Lock()
defer rm.mu.Unlock()
e, ok := rm.entries[prefix]
if !ok {
e = &allowedIPsEntry{peers: map[string]int{}}
rm.entries[prefix] = e
}
logCallerF("Increasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]",
prefix, peerKey, e.peers[peerKey], e.peers[peerKey]+1, e.total, e.total+1, e.active)
// Program WireGuard only when nothing is installed yet for this prefix.
if e.active == "" {
out, err := rm.add(prefix, peerKey)
if errors.Is(err, ErrIgnore) {
if e.total == 0 {
delete(rm.entries, prefix)
}
return Ref[string]{Count: e.total, Out: e.active}, nil
}
if err != nil {
if e.total == 0 {
delete(rm.entries, prefix)
}
return Ref[string]{}, fmt.Errorf("failed to add allowed IP %v for peer %s: %w", prefix, peerKey, err)
}
e.active = out
}
e.peers[peerKey]++
e.total++
return Ref[string]{Count: e.total, Out: e.active}, nil
}
// Decrement removes a reference to prefix for peerKey. When the peer currently installed in
// WireGuard releases its last reference, the prefix is swapped to a surviving peer if one exists,
// otherwise it is removed from WireGuard.
func (rm *AllowedIPsRefCounter) Decrement(prefix netip.Prefix, peerKey string) (Ref[string], error) {
rm.mu.Lock()
defer rm.mu.Unlock()
e, ok := rm.entries[prefix]
if !ok {
logCallerF("No allowed IP reference found for prefix %v", prefix)
return Ref[string]{}, nil
}
if e.peers[peerKey] > 0 {
logCallerF("Decreasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]",
prefix, peerKey, e.peers[peerKey], e.peers[peerKey]-1, e.total, e.total-1, e.active)
e.peers[peerKey]--
e.total--
if e.peers[peerKey] == 0 {
delete(e.peers, peerKey)
}
} else {
logCallerF("No allowed IP reference found for prefix %v peer %s", prefix, peerKey)
}
// If the peer currently installed in WireGuard still holds references, nothing to reprogram.
// Keying the check on the active peer (not the one just released) makes this self-healing:
// a prior swap whose remove/add failed leaves e.active pointing at a peer with no references,
// and this retries the hand-off on the next Decrement instead of getting stuck.
if e.active != "" && e.peers[e.active] > 0 {
return Ref[string]{Count: e.total, Out: e.active}, nil
}
// Detach the stale/gone active peer from WireGuard before reprogramming.
if e.active != "" {
if err := rm.remove(prefix, e.active); err != nil {
return Ref[string]{Count: e.total, Out: e.active}, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err)
}
e.active = ""
}
// Hand the prefix over to a surviving peer, or drop the entry when none remain.
if survivor, ok := pickSurvivor(e.peers); ok {
out, err := rm.add(prefix, survivor)
if err != nil {
return Ref[string]{Count: e.total, Out: ""}, fmt.Errorf("swap allowed IP %v to peer %s: %w", prefix, survivor, err)
}
e.active = out
return Ref[string]{Count: e.total, Out: e.active}, nil
}
delete(rm.entries, prefix)
return Ref[string]{Count: 0, Out: ""}, nil
}
// Flush removes all prefixes from WireGuard and clears the counter.
func (rm *AllowedIPsRefCounter) Flush() error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for prefix, e := range rm.entries {
if e.active == "" {
continue
}
logCallerF("Flushing allowed IP for prefix %v peer %s", prefix, e.active)
if err := rm.remove(prefix, e.active); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err))
}
}
clear(rm.entries)
return nberrors.FormatErrorOrNil(merr)
}
// ReapplyMatching calls apply for every prefix whose currently installed (active) peer satisfies
// pred, holding the lock for the whole pass. It is used to re-push allowed IPs onto a peer whose
// WireGuard entry was rebuilt (e.g. a lazy connection cycling idle->wake) without a matching
// refcounter change, which would otherwise leave the prefix installed in the counter but missing
// on the device. Only the active peer is considered — a prefix that lost its installed peer to a
// failed swap is skipped here and reconciled by the next Increment/Decrement.
func (rm *AllowedIPsRefCounter) ReapplyMatching(pred func(out string) bool, apply func(key netip.Prefix) error) error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for prefix, e := range rm.entries {
if e.active != "" && pred(e.active) {
if err := apply(prefix); err != nil {
merr = multierror.Append(merr, err)
}
}
}
return nberrors.FormatErrorOrNil(merr)
}
// pickSurvivor deterministically selects a peer still referencing the prefix. WireGuard cannot do
// multipath for a single prefix, so any surviving peer is a valid winner; the choice is made stable
// (lowest peerKey) for predictable behavior and testability.
func pickSurvivor(peers map[string]int) (string, bool) {
if len(peers) == 0 {
return "", false
}
keys := make([]string, 0, len(peers))
for k := range peers {
keys = append(keys, k)
}
sort.Strings(keys)
return keys[0], true
}

View File

@@ -1,241 +0,0 @@
package refcounter
import (
"errors"
"net/netip"
"testing"
)
// fakeWG models WireGuard's cryptokey routing: a prefix can be installed on exactly one peer.
// failAdd/failRemove make the next add/remove fail once, to exercise the self-healing error paths.
type fakeWG struct {
installed map[netip.Prefix]string
adds int
removes int
failAdd bool
failRemove bool
}
func newFakeWG() *fakeWG {
return &fakeWG{installed: map[netip.Prefix]string{}}
}
func (f *fakeWG) counter() *AllowedIPsRefCounter {
return NewAllowedIPs(
func(prefix netip.Prefix, peerKey string) (string, error) {
if f.failAdd {
f.failAdd = false
return "", errors.New("add failed")
}
f.adds++
f.installed[prefix] = peerKey
return peerKey, nil
},
func(prefix netip.Prefix, peerKey string) error {
if f.failRemove {
f.failRemove = false
return errors.New("remove failed")
}
f.removes++
// only clear if this peer is the one installed, mirroring wg semantics
if f.installed[prefix] == peerKey {
delete(f.installed, prefix)
}
return nil
},
)
}
func mustPrefix(t *testing.T, s string) netip.Prefix {
t.Helper()
p, err := netip.ParsePrefix(s)
if err != nil {
t.Fatalf("parse prefix %q: %v", s, err)
}
return p
}
func mustIncrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] {
t.Helper()
ref, err := c.Increment(p, peer)
if err != nil {
t.Fatalf("Increment(%v, %s): %v", p, peer, err)
}
return ref
}
func mustDecrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] {
t.Helper()
ref, err := c.Decrement(p, peer)
if err != nil {
t.Fatalf("Decrement(%v, %s): %v", p, peer, err)
}
return ref
}
// TestAllowedIPs_SwapOnActivePeerRemoval reproduces the reported bug: two networks with the same
// prefix routed by different peers. Removing the network whose peer is installed must hand the
// prefix over to the surviving peer instead of leaving it on the removed one.
func TestAllowedIPs_SwapOnActivePeerRemoval(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
// First peer wins while both are present.
if got := f.installed[p]; got != "peerA" {
t.Fatalf("expected peerA installed, got %q", got)
}
// Remove the active peer's network -> must swap to peerB.
mustDecrement(t, c, p, "peerA")
if got := f.installed[p]; got != "peerB" {
t.Fatalf("BUG: prefix stuck on removed peer, want peerB got %q", got)
}
// Remove the last one -> prefix gone.
mustDecrement(t, c, p, "peerB")
if _, ok := f.installed[p]; ok {
t.Fatalf("expected prefix removed, still installed on %q", f.installed[p])
}
}
// TestAllowedIPs_RemoveNonActivePeer removing a non-installed peer must not touch WireGuard.
func TestAllowedIPs_RemoveNonActivePeer(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
removesBefore := f.removes
mustDecrement(t, c, p, "peerB")
if f.installed[p] != "peerA" {
t.Fatalf("active peer must stay peerA, got %q", f.installed[p])
}
if f.removes != removesBefore {
t.Fatalf("removing a non-active peer must not call wg remove")
}
}
// TestAllowedIPs_SamePeerMultipleRefs two references via the same peer must keep the prefix until
// the last reference is released (the reason the per-peer count must be an int, not a set).
func TestAllowedIPs_SamePeerMultipleRefs(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerA")
if f.adds != 1 {
t.Fatalf("expected a single wg add for the same peer, got %d", f.adds)
}
mustDecrement(t, c, p, "peerA")
if f.installed[p] != "peerA" {
t.Fatalf("prefix must stay while a reference remains, got %q", f.installed[p])
}
if f.removes != 0 {
t.Fatalf("no wg remove expected while a reference remains, got %d", f.removes)
}
mustDecrement(t, c, p, "peerA")
if _, ok := f.installed[p]; ok {
t.Fatalf("prefix must be removed after last reference")
}
}
// TestAllowedIPs_RefCountAndActive checks the Ref returned to callers (used for the HA-disabled log).
func TestAllowedIPs_RefCountAndActive(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
ref := mustIncrement(t, c, p, "peerA")
if ref.Count != 1 || ref.Out != "peerA" {
t.Fatalf("want {1, peerA}, got {%d, %q}", ref.Count, ref.Out)
}
ref = mustIncrement(t, c, p, "peerB")
if ref.Count != 2 || ref.Out != "peerA" {
t.Fatalf("want {2, peerA}, got {%d, %q}", ref.Count, ref.Out)
}
}
// TestAllowedIPs_Flush removes everything installed and clears the counter.
func TestAllowedIPs_Flush(t *testing.T) {
f := newFakeWG()
c := f.counter()
p1 := mustPrefix(t, "10.44.8.0/24")
p2 := mustPrefix(t, "10.44.9.0/24")
mustIncrement(t, c, p1, "peerA")
mustIncrement(t, c, p2, "peerB")
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(f.installed) != 0 {
t.Fatalf("expected all prefixes removed, got %v", f.installed)
}
// After flush, a fresh increment must add again.
mustIncrement(t, c, p1, "peerC")
if f.installed[p1] != "peerC" {
t.Fatalf("counter not reset after flush")
}
}
// TestAllowedIPs_SelfHealAfterSwapAddError ensures a failed add during a swap does not permanently
// strand the prefix: the next Decrement (or Increment) must retry and install a surviving peer.
func TestAllowedIPs_SelfHealAfterSwapAddError(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
mustIncrement(t, c, p, "peerC")
// Removing the active peerA triggers a swap to a survivor; make the add fail once.
f.failAdd = true
if _, err := c.Decrement(p, "peerA"); err == nil {
t.Fatalf("expected error from failed swap add")
}
if _, ok := f.installed[p]; ok {
t.Fatalf("nothing should be installed after a failed swap add, got %q", f.installed[p])
}
// A later Decrement of a non-active survivor must retry the hand-off (self-heal), not stay stuck.
ref := mustDecrement(t, c, p, "peerC")
if got := f.installed[p]; got == "" {
t.Fatalf("self-heal failed: prefix left unrouted after add recovered")
}
if ref.Out == "" {
t.Fatalf("expected an active peer after self-heal, got empty")
}
}
// TestAllowedIPs_SelfHealAfterRemoveError ensures a failed remove during a swap is retried instead
// of leaving e.active stuck on a peer that no longer holds references.
func TestAllowedIPs_SelfHealAfterRemoveError(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
// Releasing active peerA must detach it (remove) then add peerB; fail the remove once.
f.failRemove = true
if _, err := c.Decrement(p, "peerA"); err == nil {
t.Fatalf("expected error from failed remove")
}
// Next Decrement of the non-active survivor retries: removes stale peerA, installs peerB.
mustDecrement(t, c, p, "peerB")
// peerB had only one ref, so after retry the prefix is fully released.
if _, ok := f.installed[p]; ok {
t.Fatalf("expected prefix released after self-heal, still on %q", f.installed[p])
}
}

View File

@@ -94,26 +94,6 @@ func (rm *Counter[Key, I, O]) Get(key Key) (Ref[O], bool) {
return ref, ok
}
// ReapplyMatching calls apply for every key whose stored Out satisfies pred, holding the
// counter lock for the whole pass. Running apply under the lock keeps it atomic with respect
// to Increment/Decrement: a prefix dropped to zero is removed from the map (and had its
// RemoveFunc called) before this pass observes it, so a stale key can never be re-applied.
// pred and apply are invoked under the lock, so they must not call back into the counter.
func (rm *Counter[Key, I, O]) ReapplyMatching(pred func(out O) bool, apply func(key Key) error) error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for key, ref := range rm.refCountMap {
if pred(ref.Out) {
if err := apply(key); err != nil {
merr = multierror.Append(merr, err)
}
}
}
return nberrors.FormatErrorOrNil(merr)
}
// Increment increments the reference count for the given key.
// If this is the first reference to the key, the AddFunc is called.
func (rm *Counter[Key, I, O]) Increment(key Key, in I) (Ref[O], error) {

View File

@@ -1,47 +0,0 @@
package refcounter
import (
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// TestReapplyMatching verifies ReapplyMatching invokes apply for exactly the keys whose stored
// Out satisfies the predicate (no duplicates for multiply-referenced keys) — the primitive
// ReconcilePeerAllowedIPs relies on to re-apply a single peer's routed prefixes.
func TestReapplyMatching(t *testing.T) {
rc := New[netip.Prefix, string, string](
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
func(netip.Prefix, string) error { return nil },
)
peerA1 := netip.MustParsePrefix("10.0.0.0/24")
peerA2 := netip.MustParsePrefix("10.1.0.0/24")
peerB1 := netip.MustParsePrefix("10.2.0.0/24")
for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} {
_, err := rc.Increment(prefix, peer)
require.NoError(t, err)
}
// a second reference must not make the key applied twice
_, err := rc.Increment(peerA1, "peerA")
require.NoError(t, err)
var applied []netip.Prefix
err = rc.ReapplyMatching(
func(out string) bool { return out == "peerA" },
func(key netip.Prefix) error { applied = append(applied, key); return nil },
)
require.NoError(t, err)
assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, applied)
var none []netip.Prefix
err = rc.ReapplyMatching(
func(out string) bool { return out == "missing" },
func(key netip.Prefix) error { none = append(none, key); return nil },
)
require.NoError(t, err)
assert.Empty(t, none)
}

View File

@@ -5,7 +5,5 @@ import "net/netip"
// RouteRefCounter is a Counter for Route, it doesn't take any input on Increment and doesn't use any output on Decrement
type RouteRefCounter = Counter[netip.Prefix, struct{}, struct{}]
// AllowedIPsRefCounter tracks WireGuard AllowedIPs per prefix. Unlike the generic Counter it is peer-aware:
// a prefix can be claimed by several peers at once and WireGuard allows a given prefix on exactly one peer,
// so the counter records the per-peer reference count and swaps the installed peer when the active one is released.
// See allowedips.go.
// AllowedIPsRefCounter is a Counter for AllowedIPs, it takes a peer key on Increment and passes it back to Decrement
type AllowedIPsRefCounter = Counter[netip.Prefix, string, string]

View File

@@ -15,11 +15,6 @@ type Route struct {
route *route.Route
routeRefCounter *refcounter.RouteRefCounter
allowedIPsRefcounter *refcounter.AllowedIPsRefCounter
// currentPeerKey is the routing peer this watcher currently has the prefix installed on
// (the HA winner elected by the watcher). It can differ from route.Peer and change on
// failover, so it is recorded on AddAllowedIPs and used on RemoveAllowedIPs to decrement
// the exact peer that was incremented.
currentPeerKey string
}
func NewRoute(params common.HandlerParams) *Route {
@@ -57,15 +52,12 @@ func (r *Route) AddAllowedIPs(peerKey string) error {
ref.Out,
)
}
r.currentPeerKey = peerKey
return nil
}
func (r *Route) RemoveAllowedIPs() error {
var err error
if _, decErr := r.allowedIPsRefcounter.Decrement(r.route.Network, r.currentPeerKey); decErr != nil {
err = fmt.Errorf("remove allowed IP %s: %w", r.route.Network, decErr)
if _, err := r.allowedIPsRefcounter.Decrement(r.route.Network); err != nil {
return err
}
r.currentPeerKey = ""
return err
return nil
}

View File

@@ -1,12 +0,0 @@
//go:build ios
package NetBirdSDK
import "github.com/netbirdio/netbird/version"
// GoClientVersion returns the NetBird Go client version that was baked into
// the framework at compile time via
// -ldflags "-X github.com/netbirdio/netbird/version.version=<version>".
func GoClientVersion() string {
return version.NetbirdVersion()
}

View File

@@ -1,27 +1,20 @@
package server
import (
"context"
"time"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/proto"
)
const statusCoalesceWindow = 200 * time.Millisecond
// SubscribeStatus pushes a fresh StatusResponse on every connection state
// change. The first message is the current snapshot, so a re-subscribing
// client doesn't need to also call Status. Subsequent messages fire when
// the peer recorder reports any of: connected/disconnected/connecting,
// management or signal flip, address change, or peers list change.
//
// Bursts are coalesced deterministically: the first tick after a quiet
// period is sent immediately, then a short window swallows the rest of the
// burst and a single trailing snapshot covers whatever arrived meanwhile.
// Every send is a full snapshot of the recorder's current state, so
// swallowed ticks lose no information.
// The change channel coalesces bursts to a single tick. If the consumer
// is slow the daemon drops extras (not blocks), and the next snapshot
// the consumer pulls already reflects everything.
func (s *Server) SubscribeStatus(req *proto.StatusRequest, stream proto.DaemonService_SubscribeStatusServer) error {
subID, ch := s.statusRecorder.SubscribeToStateChanges()
defer func() {
@@ -44,50 +37,12 @@ func (s *Server) SubscribeStatus(req *proto.StatusRequest, stream proto.DaemonSe
if err := s.sendStatusSnapshot(req, stream); err != nil {
return err
}
pending, open := collectStatusBurst(stream.Context(), ch)
if pending {
if err := s.sendStatusSnapshot(req, stream); err != nil {
return err
}
}
if !open {
return nil
}
case <-stream.Context().Done():
return nil
}
}
}
// collectStatusBurst waits out the coalesce window, absorbing further ticks.
// pending reports whether any tick arrived; open is false when the channel
// closed or the stream context ended.
func collectStatusBurst(ctx context.Context, ch <-chan struct{}) (pending, open bool) {
timer := time.NewTimer(statusCoalesceWindow)
defer timer.Stop()
for {
select {
case _, ok := <-ch:
if !ok {
return pending, false
}
pending = true
case <-timer.C:
select {
case _, ok := <-ch:
if !ok {
return pending, false
}
pending = true
default:
}
return pending, true
case <-ctx.Done():
return false, false
}
}
}
func (s *Server) sendStatusSnapshot(req *proto.StatusRequest, stream proto.DaemonService_SubscribeStatusServer) error {
resp, err := s.buildStatusResponse(stream.Context(), req)
if err != nil {

View File

@@ -2,7 +2,6 @@ import { type ReactNode } from "react";
import { useTranslation } from "react-i18next";
import { Browser } from "@wailsio/runtime";
import { DownloadIcon, NotepadText } from "lucide-react";
import { Update as UpdateSvc } from "@bindings/services";
import { Button } from "@/components/buttons/Button";
import { useClientVersion } from "@/contexts/ClientVersionContext";
import { cn } from "@/lib/cn";
@@ -15,12 +14,6 @@ function openUrl(url: string) {
});
}
function openInstallerDownload() {
UpdateSvc.DownloadURL()
.then(openUrl)
.catch(() => openUrl(GITHUB_RELEASES));
}
export function UpdateVersionCard() {
const { t } = useTranslation();
const { updateVersion, enforced, triggerUpdate } = useClientVersion();
@@ -44,7 +37,11 @@ export function UpdateVersionCard() {
{t("update.card.installNow")}
</Button>
) : (
<Button variant={"primary"} size={"xs"} onClick={openInstallerDownload}>
<Button
variant={"primary"}
size={"xs"}
onClick={() => openUrl(GITHUB_RELEASES)}
>
<DownloadIcon size={14} />
{t("update.card.getInstaller")}
</Button>

View File

@@ -281,9 +281,6 @@ func newApplication(onSecondInstance func()) *application.App {
Linux: application.LinuxOptions{
ProgramName: "netbird",
},
Windows: application.WindowsOptions{
WndProcInterceptor: endSessionInterceptor(),
},
SingleInstance: &application.SingleInstanceOptions{
UniqueID: "io.netbird.ui",
OnSecondInstanceLaunch: func(_ application.SecondInstanceData) {
@@ -370,9 +367,6 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat
// Hide instead of quit on close; "really quit" is reached via tray -> Quit.
window.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) {
if services.ShuttingDown() {
return
}
e.Cancel()
window.Hide()
})

View File

@@ -1,24 +0,0 @@
package services
import "sync/atomic"
var (
sessionEnding atomic.Bool
quitting atomic.Bool
)
func BeginSessionEnd() {
sessionEnding.Store(true)
}
func AbortSessionEnd() {
sessionEnding.Store(false)
}
func BeginShutdown() {
quitting.Store(true)
}
func ShuttingDown() bool {
return sessionEnding.Load() || quitting.Load()
}

View File

@@ -10,7 +10,6 @@ import (
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/ui/updater"
"github.com/netbirdio/netbird/version"
)
// UpdateResult mirrors TriggerUpdateResponse.
@@ -34,12 +33,6 @@ func (s *Update) GetState() updater.State {
return s.holder.Get()
}
// DownloadURL returns the platform-appropriate installer download link for
// manual (non-enforced) updates.
func (s *Update) DownloadURL() string {
return version.DownloadUrl()
}
// Quit exits the app. Scheduled off the calling goroutine so the JS caller's
// response returns before the runtime tears down.
func (s *Update) Quit() {

View File

@@ -154,9 +154,6 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo
})
// Hide (not destroy) on close to keep React state; reset to General for a flash-free reopen.
s.settings.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) {
if ShuttingDown() {
return
}
e.Cancel()
s.app.Event.Emit(EventSettingsOpen, "general")
s.settings.Hide()
@@ -396,9 +393,6 @@ func (s *WindowManager) CloseWelcome() {
// OpenError shows the custom error dialog; title/message are pre-localised and ride in the
// start URL. A second error replaces the open one via SetURL. Singleton, destroyed on close.
func (s *WindowManager) OpenError(title, message string) {
if ShuttingDown() {
return
}
s.mu.Lock()
defer s.mu.Unlock()
startURL := errorDialogURL(title, message)

View File

@@ -1,7 +0,0 @@
//go:build !windows && !android && !ios && !freebsd && !js
package main
func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) {
return nil
}

View File

@@ -1,36 +0,0 @@
//go:build windows
package main
import (
"os"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/ui/services"
)
const (
wmQueryEndSession = 0x0011
wmEndSession = 0x0016
)
func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) {
return func(_ uintptr, msg uint32, wParam, _ uintptr) (uintptr, bool) {
switch msg {
case wmQueryEndSession:
services.BeginSessionEnd()
return 1, true
case wmEndSession:
if wParam == 0 {
services.AbortSessionEnd()
return 0, true
}
log.Info("windows session is ending; exiting immediately")
os.Exit(0)
return 0, true
default:
return 0, false
}
}
}

View File

@@ -32,8 +32,9 @@ const (
quitDownTimeout = 5 * time.Second
urlGitHubRepo = "https://github.com/netbirdio/netbird"
urlDocs = "https://docs.netbird.io"
urlGitHubRepo = "https://github.com/netbirdio/netbird"
urlGitHubReleases = "https://github.com/netbirdio/netbird/releases/latest"
urlDocs = "https://docs.netbird.io"
)
// TrayServices bundles the services the tray menu needs, grouped so NewTray
@@ -452,7 +453,6 @@ func (t *Tray) buildMenu() *application.Menu {
}
func (t *Tray) handleQuit() {
services.BeginShutdown()
t.profileMu.Lock()
if t.switchCancel != nil {
t.switchCancel()

View File

@@ -25,9 +25,6 @@ type sendFn func(notifications.NotificationOptions) error
// event-dispatch goroutine that panic is fatal process-wide; recover() turns
// it into a logged no-op.
func safeSendNotification(send sendFn, what string, opts notifications.NotificationOptions) (err error) {
if services.ShuttingDown() {
return nil
}
defer func() {
if r := recover(); r != nil {
log.Errorf("notify %s: recovered from panic (notification bus unavailable): %v", what, r)

View File

@@ -13,7 +13,6 @@ import (
"github.com/netbirdio/netbird/client/ui/services"
"github.com/netbirdio/netbird/client/ui/updater"
"github.com/netbirdio/netbird/version"
)
// trayUpdater owns the tray UI that reacts to auto-update. Composed inside Tray.
@@ -77,15 +76,15 @@ func (u *trayUpdater) applyLanguage() {
u.refreshMenuItem(state)
}
// handleClick opens the installer download link when not Enforced, otherwise
// shows the progress page and asks the daemon to start the installer.
// handleClick opens the GitHub releases page when not Enforced, otherwise shows
// the progress page and asks the daemon to start the installer.
func (u *trayUpdater) handleClick() {
u.mu.Lock()
state := u.state
u.mu.Unlock()
if !state.Enforced {
_ = u.app.Browser.OpenURL(version.DownloadUrl())
_ = u.app.Browser.OpenURL(urlGitHubReleases)
return
}

View File

@@ -109,7 +109,7 @@ sequenceDiagram
Chk->>Inj: continue
Inj->>Inj: inject NetBird identity headers per provider config
Inj->>Grd: continue
Grd->>Grd: enforce per-provider allowlist (fail-closed backstop)
Grd->>Grd: enforce model allowlist
Grd->>Up: forward (over WireGuard)
Up-->>Resp: response (JSON or SSE stream)
Resp->>Resp: parse usage tokens, completion
@@ -135,21 +135,6 @@ sequenceDiagram
(`redact_pii = settings.RedactPii`). Phones, emails, credit cards,
PII names — see `redact.go` for the full set. See
[`modules/31-proxy-middleware-builtin.md`](modules/31-proxy-middleware-builtin.md).
- The model allowlist is enforced in TWO places. `CheckLLMPolicyLimits`
is authoritative: it resolves the policy that governs this
(provider, caller-groups) and denies (`deny_code = llm_policy.model_blocked`)
when no applicable policy permits the model — so an allowlist scoped to
one group/provider never leaks to another, and an un-guardrailed policy
is genuinely unrestricted. `llm_guardrail` is a per-provider fail-closed
backstop: it only carries an allowlist for a provider every authorising
policy restricts, and blocks unknown/undetermined models even when
management is unreachable. Because that backstop allowlist is the UNION
of every restricting policy's models, per-group narrowing lives only in
the authoritative check: during a `CheckLLMPolicyLimits` outage
`llm_limit_check` fails open, so a caller can reach any model in the
provider's union — a group scoped to model A could reach model B if
another group restricts the same provider to B. This is the documented
fail-open trade-off; a future flag may switch it to fail-closed.
- SSE streaming requires special handling on the response side; the
parser must handle partial chunks without buffering the whole
stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.md).

View File

@@ -122,7 +122,7 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest
| on_request | 1 | `llm_router` | `{"providers":[{id, models[], upstream_*, auth_header_*, allowed_group_ids[]}]}` | **true** |
| on_request | 2 | `llm_limit_check` | `{}` | |
| on_request | 3 | `llm_identity_inject` | `{"providers":[{provider_id, header_pair?, json_metadata?, extra_headers?}]}` | **true** |
| on_request | 4 | `llm_guardrail` | `{"provider_allowlists"?: {providerID: []model}, "prompt_capture":{enabled,redact_pii}}` | |
| on_request | 4 | `llm_guardrail` | `{"model_allowlist"?, "prompt_capture":{enabled,redact_pii}}` | |
| on_response | 5 | `llm_limit_record` | `{}` (runs LAST at runtime) | |
| on_response | 6 | `cost_meter` | `{}` | |
| on_response | 7 | `llm_response_parser` | `{"capture_completion": <bool>, "redact_pii"?: true}` | |

View File

@@ -244,7 +244,7 @@ no mocks. Tests: `TestChain_AllowPath_StampsAttributionAndRecordsCounter`
| `llm_router` | `{providers: [{id, models, upstream_scheme, upstream_host, upstream_path?, auth_header_name, auth_header_value, allowed_group_ids}]}` |
| `llm_limit_check` | `{}` — pulls `MgmtClient` from `FactoryContext` |
| `llm_identity_inject` | `{providers: [{provider_id, header_pair?|json_metadata?, extra_headers?}]}` |
| `llm_guardrail` | `{provider_allowlists: {providerID: []string}, prompt_capture: {enabled, redact_pii}}` — allowlist keyed by resolved provider id; a provider absent from the map is unrestricted (fail-closed backstop; authoritative per-policy/group check is management's `CheckLLMPolicyLimits`) |
| `llm_guardrail` | `{model_allowlist: []string, prompt_capture: {enabled, redact_pii}}` |
| `llm_response_parser` | `{redact_pii?, capture_completion?: *bool}` |
| `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) |
| `llm_limit_record` | `{}` — same pattern as `llm_limit_check` |

View File

@@ -9,234 +9,25 @@ import (
"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"
)
// per1k is a model's published USD rates per 1k tokens. read is the prompt-cache read rate
// (OpenAI: the cached-input discount rate); write is the cache-creation rate where one exists.
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.
var publishedPer1k = map[string]per1k{
"gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0},
"gpt-4o": {0.0025, 0.01, 0.00125, 0},
"claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125},
"claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375},
"claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375},
"kimi-k3": {0.003, 0.015, 0.0003, 0.003}, // no published write rate: bills at the input rate
"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},
}
// 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.
const rawCostVerificationSQL = `
WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS (
VALUES
('gpt-4o-mini', 0.00015, 0.0006, 0.000075, 0.0),
('gpt-4o', 0.0025, 0.01, 0.00125, 0.0),
('claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125),
('claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375),
('claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375),
('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)
)
SELECT
u.provider,
u.model,
u.input_tokens,
u.output_tokens,
u.cached_input_tokens,
u.cache_creation_tokens,
u.input_cost_usd,
u.cached_input_cost_usd,
u.cache_creation_cost_usd,
u.output_cost_usd,
-- No cost_usd / cache_cost_usd columns are stored: both are derived from the
-- four per-bucket columns above, exactly as the API renders them.
(u.input_cost_usd + u.cached_input_cost_usd + u.cache_creation_cost_usd + u.output_cost_usd) AS cost_usd,
(u.cached_input_cost_usd + u.cache_creation_cost_usd) AS cache_cost_usd,
CASE WHEN u.provider = 'openai' THEN
(u.input_tokens - MIN(u.cached_input_tokens, u.input_tokens))*r.in_rate/1000.0
ELSE
u.input_tokens*r.in_rate/1000.0
END AS expected_input,
CASE WHEN u.provider = 'openai' THEN
MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
ELSE
u.cached_input_tokens*r.read_rate/1000.0
END AS expected_cached_input,
CASE WHEN u.provider = 'openai' THEN
0.0
ELSE
u.cache_creation_tokens*r.write_rate/1000.0
END AS expected_cache_creation,
u.output_tokens*r.out_rate/1000.0 AS expected_output,
CASE WHEN u.provider = 'openai' THEN
(u.input_tokens - MIN(u.cached_input_tokens, u.input_tokens))*r.in_rate/1000.0
+ MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
+ u.output_tokens*r.out_rate/1000.0
ELSE
u.input_tokens*r.in_rate/1000.0 + u.cached_input_tokens*r.read_rate/1000.0
+ u.cache_creation_tokens*r.write_rate/1000.0 + u.output_tokens*r.out_rate/1000.0
END AS expected_total,
CASE WHEN u.provider = 'openai' THEN
MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
ELSE
u.cached_input_tokens*r.read_rate/1000.0 + u.cache_creation_tokens*r.write_rate/1000.0
END AS expected_cache
FROM agent_network_request_usage u
JOIN rates r ON r.model = u.model
ORDER BY u.timestamp`
// verifyUsageRowsSQL re-checks every persisted usage row directly in the management sqlite store,
// bypassing the API path — the same audit an operator can run on a production store.db.
func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) {
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() }()
rows, err := db.Raw(rawCostVerificationSQL).Rows()
require.NoError(t, err, "run raw cost verification query")
defer func() { _ = rows.Close() }()
verified := 0
for rows.Next() {
var provider, model string
var inTok, outTok, readTok, writeTok int64
var inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost float64
var wantInput, wantCachedInput, wantCacheCreation, wantOutput, wantTotal, wantCache float64
require.NoError(t, rows.Scan(&provider, &model, &inTok, &outTok, &readTok, &writeTok,
&inCost, &cachedInCost, &cacheCreateCost, &outCost, &cost, &cacheCost,
&wantInput, &wantCachedInput, &wantCacheCreation, &wantOutput, &wantTotal, &wantCache), "scan usage row")
t.Logf("[sql] %s/%s: in=%d out=%d cache_read=%d cache_write=%d stored in/cached/create/out=$%.6f/$%.6f/$%.6f/$%.6f total=$%.6f cache=$%.6f expected total=$%.6f cache=$%.6f",
provider, model, inTok, outTok, readTok, writeTok,
inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost, wantTotal, wantCache)
assert.InDeltaf(t, wantInput, inCost, 1e-6, "stored input_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCachedInput, cachedInCost, 1e-6, "stored cached_input_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCacheCreation, cacheCreateCost, 1e-6, "stored cache_creation_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantOutput, outCost, 1e-6, "stored output_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantTotal, cost, 1e-6, "derived cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCache, cacheCost, 1e-6, "derived cache_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, inCost+cachedInCost+cacheCreateCost+outCost, cost, 1e-9,
"stored buckets must sum to the derived cost_usd for %s/%s", provider, model)
verified++
}
require.NoError(t, rows.Err(), "iterate usage rows")
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)
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()
require.NoError(t, err, "query gateway-prefixed usage rows")
defer func() { _ = gwRows.Close() }()
for gwRows.Next() {
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)
}
require.NoError(t, gwRows.Err(), "iterate gateway usage rows")
}
// validateAccessLogCost recomputes a live access-log row's expected total and cache cost from the
// published per-1k rates and the row's persisted token buckets, and asserts both stored values.
// Gateway-prefixed model ids the proxy deliberately does not price must store cost 0.
func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAccessLog) {
t.Helper()
model := catalogModel(pc)
provider := ""
if row.Provider != nil {
provider = *row.Provider
}
t.Logf("[cost] %s: provider=%s model=%s in=%d out=%d total=%d cache_read=%d cache_write=%d cost=$%.6f cache_cost=$%.6f",
pc.name, provider, model, row.InputTokens, row.OutputTokens, row.TotalTokens,
row.CachedInputTokens, row.CacheCreationTokens, row.CostUsd, row.CacheCostUsd)
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
}
// input_tokens may legitimately be 0: Moonshot/Kimi reports fully cached prompts under the cache
// buckets only. Output and total must always be present on a priced row.
require.Positive(t, row.OutputTokens, "priced row must carry output tokens")
require.Positive(t, row.TotalTokens, "priced row must carry total tokens")
var wantInput, wantCachedInput, wantCacheCreation float64
if provider == "openai" {
cached := min(row.CachedInputTokens, row.InputTokens) // cached is a subset of input
wantInput = float64(row.InputTokens-cached) / 1000 * rates.in
wantCachedInput = float64(cached) / 1000 * rates.read
// OpenAI has no cache-write bucket; wantCacheCreation stays 0.
} else {
// Anthropic / Bedrock shape: cache buckets are additive to input_tokens.
wantInput = float64(row.InputTokens) / 1000 * rates.in
wantCachedInput = float64(row.CachedInputTokens) / 1000 * rates.read
wantCacheCreation = float64(row.CacheCreationTokens) / 1000 * rates.write
}
wantOutput := float64(row.OutputTokens) / 1000 * rates.out
wantCache := wantCachedInput + wantCacheCreation
wantTotal := wantInput + wantCache + wantOutput
t.Logf("[cost] %s: expecting input=$%.6f cached_input=$%.6f cache_creation=$%.6f output=$%.6f total=$%.6f cache=$%.6f from published rates",
pc.name, wantInput, wantCachedInput, wantCacheCreation, wantOutput, wantTotal, wantCache)
assert.InDeltaf(t, wantInput, row.InputCostUsd, 1e-6, "stored input_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCachedInput, row.CachedInputCostUsd, 1e-6, "stored cached_input_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCacheCreation, row.CacheCreationCostUsd, 1e-6, "stored cache_creation_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantOutput, row.OutputCostUsd, 1e-6, "stored output_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "derived cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCache, row.CacheCostUsd, 1e-6, "derived cache_cost_usd for %s (%s)", pc.name, model)
// The aggregates must be exactly the sum of the stored components, not an
// independently-computed figure that could drift from the breakdown.
assert.InDeltaf(t, row.InputCostUsd+row.CachedInputCostUsd+row.CacheCreationCostUsd+row.OutputCostUsd,
row.CostUsd, 1e-9, "stored buckets must sum to the derived cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, row.CachedInputCostUsd+row.CacheCreationCostUsd,
row.CacheCostUsd, 1e-9, "stored cache buckets must sum to the derived cache_cost_usd for %s (%s)", pc.name, model)
}
// providerCase is one entry in the live provider matrix. The same scenario runs
// for every available provider; availability is keyed off env vars so the suite
// covers whatever credentials are present (source ~/.llm-keys locally / set the
// Actions secrets in CI).
type providerCase struct {
name string
catalogID string
upstream string
apiKey string
model string // body model (chat/messages) or path model@version (vertex)
kind string // harness.WireChat, harness.WireMessages, or harness.WireVertex
project string // vertex only: GCP project for the rawPredict path
region string // vertex only: GCP region for the rawPredict path
pathPrefix string // base-URL path prefix the agent carries (e.g. "/anthropic" for Kimi)
name string
catalogID string
upstream string
apiKey string
model string // body model (chat/messages) or path model@version (vertex)
kind string // harness.WireChat, harness.WireMessages, or harness.WireVertex
project string // vertex only: GCP project for the rawPredict path
region string // vertex only: GCP region for the rawPredict path
}
// availableProviders builds the matrix from the provider env vars that are set.
@@ -248,24 +39,6 @@ func availableProviders() []providerCase {
if k := os.Getenv("ANTHROPIC_TOKEN"); k != "" {
ps = append(ps, providerCase{name: "anthropic", catalogID: "anthropic_api", upstream: "https://api.anthropic.com", apiKey: k, model: "claude-haiku-4-5", kind: harness.WireMessages})
}
if k := os.Getenv("KIMI_TOKEN"); k != "" {
// Kimi (Moonshot AI) serves two body shapes from the same key: OpenAI
// Chat Completions on the bare host (/v1/...) and the Anthropic
// Messages API under the /anthropic path prefix (the endpoint
// Moonshot's Claude Code guide uses). The provider keeps the bare
// default upstream and the AGENT carries the /anthropic prefix in
// its base URL — exactly the documented Claude Code / Kimi CLI
// setup (ANTHROPIC_BASE_URL=https://<endpoint>/anthropic) — so one
// provider serves both shapes and the prefix rides through to
// Moonshot. Run the Anthropic shape, the flagship Claude Code path;
// the OpenAI wire shape is covered live by the other chat-shaped
// matrix providers, and Kimi-over-chat passed with kimi-k3 before
// the single-model constraint surfaced (run #73 on the kimi feature
// branch). The platform serves this account exactly ONE model —
// kimi-k3 (kimi-k2-thinking and even kimi-latest return
// resource_not_found_error on both surfaces).
ps = append(ps, providerCase{name: "kimi", catalogID: "kimi_api", upstream: "https://api.moonshot.ai", apiKey: k, model: "kimi-k3", kind: harness.WireMessages, pathPrefix: "/anthropic"})
}
if k, u := os.Getenv("VERCEL_TOKEN"), os.Getenv("VERCEL_URL"); k != "" && u != "" {
ps = append(ps, providerCase{name: "vercel", catalogID: "vercel_ai_gateway", upstream: u, apiKey: k, model: "openai/gpt-4o-mini", kind: harness.WireChat})
}
@@ -311,25 +84,19 @@ func availableProviders() []providerCase {
}
}
// Bedrock: path-routed, bearer auth. Model is the FULL cross-region
// inference-profile id exactly as AWS issues it — region-family prefix
// plus the date/version suffix. A bare or wrong-region id makes Bedrock
// reject the request with "The provided model identifier is invalid"
// before any inference runs. The proxy normalizes this id to the catalog
// key (anthropic.claude-haiku-4-5) for routing/pricing/allowlists.
// Defaults pair eu-central-1 with the eu.* profile; AWS_REGION overrides
// the region and the prefix follows its family.
// Bedrock: path-routed, bearer auth. Model is a cross-region inference
// profile id (distinct string from the first-party Anthropic case).
if k := os.Getenv("AWS_BEARER_TOKEN_BEDROCK"); k != "" {
region := os.Getenv("AWS_REGION")
if region == "" {
region = "eu-central-1"
region = "us-east-1"
}
// A valid Bedrock inference-profile id, overridable per account (AWS_BEDROCK_MODEL, also the
// workflow's bedrock_model dispatch input). `global.` profiles work from any region. Defaults to
// Sonnet 4.6, whose id convention dropped the -YYYYMMDD-v1:0 suffix that Haiku 4.5 still carries.
// A valid Bedrock inference-profile id (region prefix + date + version),
// overridable per account. `global.` profiles can be invoked from any
// region; set AWS_BEDROCK_MODEL to match the enabled profile for the token.
model := os.Getenv("AWS_BEDROCK_MODEL")
if model == "" {
model = "global.anthropic.claude-sonnet-4-6"
model = "global.anthropic.claude-haiku-4-5-20251001-v1:0"
}
ps = append(ps, providerCase{name: "bedrock", catalogID: "bedrock_api", upstream: "https://bedrock-runtime." + region + ".amazonaws.com", apiKey: k, model: model, kind: harness.WireBedrock})
}
@@ -449,7 +216,8 @@ func TestProvidersMatrix(t *testing.T) {
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 (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the proxy peer so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve agent-network endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
@@ -465,10 +233,6 @@ func TestProvidersMatrix(t *testing.T) {
// session id and confirm the marker propagated end-to-end.
sessionID := "e2e-session-" + pc.name
// A long-form prompt so completions carry realistic token counts for cost validation;
// max_tokens in the harness bodies (2048) lets the full answer through.
const matrixPrompt = "explain GitHub workflow in 1000 words"
// Retry briefly to absorb tunnel/DNS jitter on the first call.
var code int
var body string
@@ -479,11 +243,11 @@ func TestProvidersMatrix(t *testing.T) {
var cerr error
switch pc.kind {
case harness.WireVertex:
c, b, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, pc.project, pc.region, pc.model, matrixPrompt, sessionID)
c, b, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, pc.project, pc.region, pc.model, "Reply with exactly: pong", sessionID)
case harness.WireBedrock:
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, matrixPrompt, sessionID)
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, "Reply with exactly: pong", sessionID)
default:
c, b, cerr = cl.ChatPrefixed(ctx, settings.Endpoint, proxyIP, pc.pathPrefix, pc.kind, pc.model, matrixPrompt, sessionID)
c, b, cerr = cl.Chat(ctx, settings.Endpoint, proxyIP, pc.kind, pc.model, "Reply with exactly: pong", sessionID)
}
if cerr == nil {
code, body = c, b
@@ -502,7 +266,6 @@ func TestProvidersMatrix(t *testing.T) {
// The session id sent as x-session-id must round-trip into the
// access-log row for this provider.
var row api.AgentNetworkAccessLog
require.Eventually(t, func() bool {
logs, lerr := srv.ListAccessLogs(ctx)
if lerr != nil {
@@ -510,15 +273,11 @@ func TestProvidersMatrix(t *testing.T) {
}
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 for %s", sessionID, pc.name)
// Stored total and cache cost must match the published rates applied to the row's buckets.
validateAccessLogCost(t, pc, row)
})
}
@@ -539,7 +298,4 @@ func TestProvidersMatrix(t *testing.T) {
}
return false
}, 60*time.Second, 3*time.Second, "consumption must be recorded with positive token counts after live traffic")
// Final raw-SQL audit: bypass the API and re-verify every persisted usage row in the store.
verifyUsageRowsSQL(t, srv)
}

View File

@@ -1,209 +0,0 @@
//go:build e2e
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// pathRoutedGuardrailCase is one provider's self-contained scenario: its own
// provider, its own guardrail whose allowlist holds ONLY that provider's
// allowed model, and its own policy. Each case runs in isolation (its own
// proxy + client), so the guardrail the proxy enforces contains exactly this
// provider's model — never a mixed cross-provider list.
type pathRoutedGuardrailCase struct {
name string
catalogID string // agent-network catalog provider id
wire string // harness.WireVertex | harness.WireBedrock
allowEntry string // the single model id put on the guardrail allowlist
allowModel string // model id sent that MUST be served (200)
blockModel string // model id sent that MUST be denied (403 model_blocked)
}
// TestGuardrailBlocksUnselectedModel_PathRouted is the end-to-end regression
// guard for the customer report that a model-allowlist guardrail attached to a
// policy has no effect for PATH-ROUTED providers — where the model travels in
// the URL, not the JSON body: Google Vertex (…/models/{model}:rawPredict) and
// AWS Bedrock (/model/{id}/invoke).
//
// Each provider is tested in isolation with a guardrail allowlisting a single
// model of its own: the allowed model (in the URL path) is served (200) and an
// unselected model (in the URL path) is denied 403 by the guardrail
// (llm_policy.model_blocked) before the upstream. The Vertex case mirrors the
// customer verbatim — allow Sonnet, and the unselected model is the exact
// claude-opus-4-6 they reported reaching the model unblocked. The Bedrock case
// sends a region-prefixed, versioned inference-profile id so URL-path model
// normalization is exercised too.
//
// The provider is catch-all (no models), so the router forwards any model and a
// 403 can only come from the guardrail, never model_not_routable. Only the
// upstream LLM is mocked (the vLLM nginx answers any path with 200); management
// synth/reconcile, the proxy middleware chain (URL-path model extraction,
// router, guardrail) and the tunnel are all real, and the guardrail denies
// before the upstream is dialed so the mock cannot influence the block. A
// static bearer api key is used so the router injects a static Authorization
// header instead of minting a GCP token — the only reason path-routed providers
// normally need live credentials — so the test runs with none and is always on.
func TestGuardrailBlocksUnselectedModel_PathRouted(t *testing.T) {
cases := []pathRoutedGuardrailCase{
{
name: "vertex",
catalogID: "vertex_ai_api",
wire: harness.WireVertex,
allowEntry: "claude-sonnet-4-5",
allowModel: "claude-sonnet-4-5",
blockModel: "claude-opus-4-6", // the customer-reported model
},
{
name: "bedrock",
catalogID: "bedrock_api",
wire: harness.WireBedrock,
allowEntry: "anthropic.claude-sonnet-4-5", // normalized catalog id
allowModel: "us.anthropic.claude-sonnet-4-5-v1:0",
blockModel: "us.anthropic.claude-opus-4-8-v1:0",
},
}
for _, tc := range cases {
tc := tc
t.Run(tc.name, func(t *testing.T) {
runPathRoutedGuardrailCase(t, tc)
})
}
}
func runPathRoutedGuardrailCase(t *testing.T, tc pathRoutedGuardrailCase) {
t.Helper()
const (
vertexProject = "e2e-project"
vertexRegion = "global"
)
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
defer cancel()
vllm, err := harness.StartVLLM(ctx, srv)
require.NoError(t, err, "start mock upstream")
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-" + tc.name})
require.NoError(t, err, "create group")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grp.Id) })
ephemeral := false
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
Name: "e2e-guardrail-" + tc.name + "-client",
Type: "reusable",
ExpiresIn: 86400,
UsageLimit: 0,
AutoGroups: []string{grp.Id},
Ephemeral: &ephemeral,
})
require.NoError(t, err, "mint setup key")
require.NotEmpty(t, sk.Key, "setup key plaintext")
// Catch-all provider (no models) so the router forwards any model; a static
// bearer key means the router injects a static auth header instead of minting
// a GCP token. Bootstraps the cluster if it isn't already.
staticKey := "static-e2e-token"
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
Name: tc.name,
ProviderId: tc.catalogID,
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
BootstrapCluster: ptr(harness.AgentNetworkCluster),
})
require.NoError(t, err, "create %s provider", tc.name)
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
// Guardrail allowlisting ONLY this provider's allowed model.
var gr api.AgentNetworkGuardrailRequest
gr.Name = "e2e-guardrail-" + tc.name
gr.Checks.ModelAllowlist.Enabled = true
gr.Checks.ModelAllowlist.Models = []string{tc.allowEntry}
guard, err := srv.CreateGuardrail(ctx, gr)
require.NoError(t, err, "create guardrail")
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), guard.Id) })
enabled := true
pol, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-guardrail-" + tc.name,
Enabled: &enabled,
SourceGroups: []string{grp.Id},
DestinationProviderIds: []string{prov.Id},
GuardrailIds: &[]string{guard.Id},
})
require.NoError(t, err, "create policy")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), pol.Id) })
settings, err := srv.GetSettings(ctx)
require.NoError(t, err, "read settings")
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-"+tc.name+"-proxy")
require.NoError(t, err, "mint proxy token")
px, err := harness.StartProxy(ctx, srv, proxyToken)
require.NoError(t, err, "start proxy")
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
cl, err := harness.StartClient(ctx, srv, sk.Key)
require.NoError(t, err, "start client")
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
// Probe first: the GET resolves the endpoint and its first packet wakes the
// lazy proxy peer, so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
}
send := func(model string) (int, string) {
var code int
var body string
var cerr error
switch tc.wire {
case harness.WireVertex:
code, body, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, vertexProject, vertexRegion, model, "Reply with exactly: pong", "")
case harness.WireBedrock:
code, body, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, model, "Reply with exactly: pong", "")
default:
t.Fatalf("unsupported wire %q", tc.wire)
}
require.NoError(t, cerr, "request must reach the proxy for %s", tc.name)
return code, body
}
// Allowed model (in the URL path) is served. Retry to absorb tunnel/DNS
// jitter on the first call over the freshly warmed tunnel.
var code int
var body string
deadline := time.Now().Add(90 * time.Second)
for time.Now().Before(deadline) {
code, body = send(tc.allowModel)
if code == 200 {
break
}
time.Sleep(5 * time.Second)
}
assert.Equal(t, 200, code,
"allowed %s model (URL path) must be served; body: %s\n=== proxy logs ===\n%s", tc.name, body, px.Logs(context.Background()))
// Unselected model (in the URL path) must be blocked by the guardrail.
code, body = send(tc.blockModel)
assert.Equal(t, 403, code,
"unselected %s model (URL path) must be denied, not served; body: %s\n=== proxy logs ===\n%s", tc.name, body, px.Logs(context.Background()))
assert.Contains(t, body, "llm_policy.model_blocked",
"%s denial must come from the guardrail allowlist, not routing; body: %s", tc.name, body)
}

View File

@@ -1,205 +0,0 @@
//go:build e2e
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// TestGuardrailGroupSwitchTakesEffectAfterTTL proves that moving a peer between
// groups flips its model-allowlist decision once the proxy's tunnel-peer cache
// expires. The peer's groups reach the guardrail via ValidateTunnelPeer, which
// the proxy caches; the switch is invisible until that cache expires. The proxy
// runs with a short NB_PROXY_TUNNEL_CACHE_TTL so the flip happens in seconds
// instead of the 5-minute default.
//
// Setup: one catch-all provider declaring modelA + modelB; polA (grpA -> allow
// modelA) and polB (grpB -> allow modelB). The client starts in grpA. modelA is
// served and modelB denied; after switching the client grpA -> grpB, modelB is
// served and modelA denied. The cross-group deny comes from management's
// per-policy/group CheckLLMPolicyLimits (the proxy backstop carries the union).
func TestGuardrailGroupSwitchTakesEffectAfterTTL(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
defer cancel()
const (
modelA = "e2e-model-a"
modelB = "e2e-model-b"
// Short tunnel-cache TTL so a group switch propagates in seconds.
// Exercises the NB_PROXY_TUNNEL_CACHE_TTL override.
cacheTTL = 3 * time.Second
)
vllm, err := harness.StartVLLM(ctx, srv)
require.NoError(t, err, "start mock upstream")
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
grpA, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-gswitch-a"})
require.NoError(t, err, "create group A")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpA.Id) })
grpB, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-gswitch-b"})
require.NoError(t, err, "create group B")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpB.Id) })
ephemeral := false
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
Name: "e2e-gswitch-client",
Type: "reusable",
ExpiresIn: 86400,
UsageLimit: 0,
AutoGroups: []string{grpA.Id}, // client starts in group A
Ephemeral: &ephemeral,
})
require.NoError(t, err, "mint setup key")
require.NotEmpty(t, sk.Key, "setup key plaintext")
staticKey := "static-e2e-token"
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
Name: "gswitch",
ProviderId: "openai_api",
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
Models: &[]api.AgentNetworkProviderModel{
{Id: modelA, InputPer1k: 0.001, OutputPer1k: 0.001},
{Id: modelB, InputPer1k: 0.001, OutputPer1k: 0.001},
},
BootstrapCluster: ptr(harness.AgentNetworkCluster),
})
require.NoError(t, err, "create provider")
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
mkGuard := func(name, model string) api.AgentNetworkGuardrail {
var gr api.AgentNetworkGuardrailRequest
gr.Name = name
gr.Checks.ModelAllowlist.Enabled = true
gr.Checks.ModelAllowlist.Models = []string{model}
g, gerr := srv.CreateGuardrail(ctx, gr)
require.NoError(t, gerr, "create guardrail %s", name)
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) })
return g
}
gA := mkGuard("e2e-gswitch-a", modelA)
gB := mkGuard("e2e-gswitch-b", modelB)
enabled := true
polA, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-gswitch-a",
Enabled: &enabled,
SourceGroups: []string{grpA.Id},
DestinationProviderIds: []string{prov.Id},
GuardrailIds: &[]string{gA.Id},
})
require.NoError(t, err, "create policy A")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polA.Id) })
polB, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-gswitch-b",
Enabled: &enabled,
SourceGroups: []string{grpB.Id},
DestinationProviderIds: []string{prov.Id},
GuardrailIds: &[]string{gB.Id},
})
require.NoError(t, err, "create policy B")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polB.Id) })
settings, err := srv.GetSettings(ctx)
require.NoError(t, err, "read settings")
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-gswitch-proxy")
require.NoError(t, err, "mint proxy token")
px, err := harness.StartProxy(ctx, srv, proxyToken, map[string]string{
"NB_PROXY_TUNNEL_CACHE_TTL": cacheTTL.String(),
})
require.NoError(t, err, "start proxy")
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
cl, err := harness.StartClient(ctx, srv, sk.Key)
require.NoError(t, err, "start client")
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
}
send := func(model string) (int, string) {
code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "")
require.NoError(t, cerr, "request must reach the proxy")
return code, body
}
sendUntil := func(model string, want int, timeout time.Duration) (int, string) {
var code int
var body string
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
code, body = send(model)
if code == want {
return code, body
}
time.Sleep(2 * time.Second)
}
return code, body
}
// Phase 1 — client is in group A: modelA served, modelB denied.
code, body := sendUntil(modelA, 200, 90*time.Second)
assert.Equal(t, 200, code,
"group-A model must be served while the client is in group A; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
code, body = send(modelB)
assert.Equal(t, 403, code,
"group-B model must be denied while the client is in group A; body: %s", body)
assert.Contains(t, body, "llm_policy.model_blocked")
// Switch the client peer from group A to group B.
peerID := clientPeerInGroup(t, ctx, grpA.Id)
_, err = srv.API().Groups.Update(ctx, grpB.Id, api.PutApiGroupsGroupIdJSONRequestBody{
Name: grpB.Name,
Peers: &[]string{peerID},
})
require.NoError(t, err, "add peer to group B")
_, err = srv.API().Groups.Update(ctx, grpA.Id, api.PutApiGroupsGroupIdJSONRequestBody{
Name: grpA.Name,
Peers: &[]string{},
})
require.NoError(t, err, "remove peer from group A")
// Phase 2 — after the short TTL expires the proxy re-validates the peer,
// sees group B, and the decision flips. Poll to absorb TTL + re-validation.
code, body = sendUntil(modelB, 200, 60*time.Second)
assert.Equal(t, 200, code,
"after the group switch + TTL, the group-B model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
code, body = send(modelA)
assert.Equal(t, 403, code,
"after the switch, the old group-A model must be denied; body: %s", body)
assert.Contains(t, body, "llm_policy.model_blocked")
}
// clientPeerInGroup returns the id of the single peer that is a member of the
// given group — the test client. The proxy peer is never added to test groups.
func clientPeerInGroup(t *testing.T, ctx context.Context, groupID string) string {
t.Helper()
peers, err := srv.API().Peers.List(ctx)
require.NoError(t, err, "list peers")
for _, p := range peers {
for _, g := range p.Groups {
if g.Id == groupID {
return p.Id
}
}
}
t.Fatalf("no peer found in group %s", groupID)
return ""
}

View File

@@ -1,201 +0,0 @@
//go:build e2e
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// TestGuardrailMultiPolicyModelAllowlist: modelSelected served (200), grpOther's
// modelOther denied for the grpMain client (403 model_blocked, no cross-group
// leak), and openModel on the un-guardrailed policy's provider served (200).
func TestGuardrailMultiPolicyModelAllowlist(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
defer cancel()
const (
modelSelected = "e2e-selected"
modelOther = "e2e-other"
openModel = "e2e-open"
)
vllm, err := harness.StartVLLM(ctx, srv)
require.NoError(t, err, "start mock upstream")
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-main"})
require.NoError(t, err, "create main group")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) })
grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-other"})
require.NoError(t, err, "create other group")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) })
ephemeral := false
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
Name: "e2e-guardrail-mp-client",
Type: "reusable",
ExpiresIn: 86400,
UsageLimit: 0,
AutoGroups: []string{grpMain.Id}, // client joins grpMain only
Ephemeral: &ephemeral,
})
require.NoError(t, err, "mint setup key")
require.NotEmpty(t, sk.Key, "setup key plaintext")
staticKey := "static-e2e-token"
models := func(ids ...string) *[]api.AgentNetworkProviderModel {
out := make([]api.AgentNetworkProviderModel, 0, len(ids))
for _, id := range ids {
out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001})
}
return &out
}
// pRestricted declares the two guardrailed models so routing is deterministic
// (model -> provider). Created first, so it carries the bootstrap cluster.
pRestricted, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
Name: "restricted",
ProviderId: "openai_api",
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
Models: models(modelSelected, modelOther),
BootstrapCluster: ptr(harness.AgentNetworkCluster),
})
require.NoError(t, err, "create restricted provider")
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pRestricted.Id) })
pOpen, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
Name: "open",
ProviderId: "openai_api",
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
Models: models(openModel),
})
require.NoError(t, err, "create open provider")
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pOpen.Id) })
mkGuardrail := func(name, model string) api.AgentNetworkGuardrail {
var gr api.AgentNetworkGuardrailRequest
gr.Name = name
gr.Checks.ModelAllowlist.Enabled = true
gr.Checks.ModelAllowlist.Models = []string{model}
g, gerr := srv.CreateGuardrail(ctx, gr)
require.NoError(t, gerr, "create guardrail %s", name)
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) })
return g
}
gMain := mkGuardrail("e2e-guardrail-mp-main", modelSelected)
gOther := mkGuardrail("e2e-guardrail-mp-other", modelOther)
enabled := true
// polMain: grpMain restricted to modelSelected on pRestricted.
polMain, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-guardrail-mp-main",
Enabled: &enabled,
SourceGroups: []string{grpMain.Id},
DestinationProviderIds: []string{pRestricted.Id},
GuardrailIds: &[]string{gMain.Id},
})
require.NoError(t, err, "create main policy")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) })
// polOther: grpOther restricted to modelOther on the SAME provider. The
// client is not in grpOther, so modelOther must never be usable by it.
polOther, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-guardrail-mp-other",
Enabled: &enabled,
SourceGroups: []string{grpOther.Id},
DestinationProviderIds: []string{pRestricted.Id},
GuardrailIds: &[]string{gOther.Id},
})
require.NoError(t, err, "create other policy")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) })
// polOpen: grpMain on pOpen with NO guardrail — unrestricted.
polOpen, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-guardrail-mp-open",
Enabled: &enabled,
SourceGroups: []string{grpMain.Id},
DestinationProviderIds: []string{pOpen.Id},
})
require.NoError(t, err, "create open policy")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOpen.Id) })
settings, err := srv.GetSettings(ctx)
require.NoError(t, err, "read settings")
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-mp-proxy")
require.NoError(t, err, "mint proxy token")
px, err := harness.StartProxy(ctx, srv, proxyToken)
require.NoError(t, err, "start proxy")
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
cl, err := harness.StartClient(ctx, srv, sk.Key)
require.NoError(t, err, "start client")
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
}
send := func(model string) (int, string) {
code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "")
require.NoError(t, cerr, "request must reach the proxy")
return code, body
}
// sendUntil200 absorbs first-call tunnel/DNS jitter on the freshly warmed tunnel.
sendUntil200 := func(model string) (int, string) {
var code int
var body string
deadline := time.Now().Add(90 * time.Second)
for time.Now().Before(deadline) {
code, body = send(model)
if code == 200 {
break
}
time.Sleep(5 * time.Second)
}
return code, body
}
t.Run("selected model allowed for its group", func(t *testing.T) {
code, body := sendUntil200(modelSelected)
assert.Equal(t, 200, code,
"grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
})
t.Run("other group's model does not leak", func(t *testing.T) {
// modelOther is allowlisted only for grpOther. The grpMain client must be
// denied by management's per-policy/group check — not waved through by an
// account-wide union. This is the security-critical wrong-ALLOW guard.
code, body := send(modelOther)
assert.Equal(t, 403, code,
"another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
assert.Contains(t, body, "llm_policy.model_blocked",
"denial must be a model-allowlist decision; body: %s", body)
})
t.Run("unguarded policy leaves its provider unrestricted", func(t *testing.T) {
// polOpen carries no guardrail, so pOpen is unrestricted for grpMain. The
// old account-wide union would have blocked openModel (it is on no
// allowlist); it must now be served — the false-DENY guard.
code, body := sendUntil200(openModel)
assert.Equal(t, 200, code,
"an un-guardrailed policy's provider must not be blocked by another policy's allowlist; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
})
}

View File

@@ -1,422 +0,0 @@
//go:build e2e
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// pergroupCase describes one provider surface for the per-group allowlist matrix.
// selectedReq/otherReq are the model identifiers as they travel in the request
// (URL path for Bedrock/Vertex, body "model" for chat/messages). selectedAllow/
// otherAllow are the (normalized) forms the guardrail allowlist holds — for
// Bedrock these differ from the request form so path normalization is exercised.
type pergroupCase struct {
name string
catalogID string
wire string // "chat", "messages", "vertex", "bedrock"
models *[]api.AgentNetworkProviderModel
selectedReq string
selectedAllow string
otherReq string
otherAllow string
providerID string // filled during setup
}
// TestGuardrailPerGroupAllowlist_AllProviders proves the per-policy/group model
// allowlist end to end across every always-on provider surface, including the
// path-routed ones (Vertex, Bedrock) where the model travels in the URL.
//
// For each provider two policies target it: grpMain (the client) is allowed only
// selectedReq; grpOther (which the client is NOT in) is allowed only otherReq.
// The client must get selectedReq served (200) and otherReq denied (403,
// llm_policy.model_blocked) — the cross-group no-leak property. The deny is the
// authoritative per-policy/group decision from management (the proxy per-provider
// backstop carries the union of both models), so this also confirms management
// receives the correct normalized model for path-routed providers.
func TestGuardrailPerGroupAllowlist_AllProviders(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Minute)
defer cancel()
const (
vertexProject = "e2e-project"
vertexRegion = "global"
)
priced := func(ids ...string) *[]api.AgentNetworkProviderModel {
out := make([]api.AgentNetworkProviderModel, 0, len(ids))
for _, id := range ids {
out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001})
}
return &out
}
cases := []*pergroupCase{
{
name: "openai", catalogID: "openai_api", wire: harness.WireChat,
models: priced("oai-model-a", "oai-model-b"),
selectedReq: "oai-model-a", selectedAllow: "oai-model-a",
otherReq: "oai-model-b", otherAllow: "oai-model-b",
},
{
name: "anthropic", catalogID: "anthropic_api", wire: harness.WireMessages,
models: priced("ant-model-a", "ant-model-b"),
selectedReq: "ant-model-a", selectedAllow: "ant-model-a",
otherReq: "ant-model-b", otherAllow: "ant-model-b",
},
{
// Vertex catalog ids travel bare in the rawPredict path.
name: "vertex", catalogID: "vertex_ai_api", wire: "vertex",
selectedReq: "claude-sonnet-4-5", selectedAllow: "claude-sonnet-4-5",
otherReq: "claude-opus-4-6", otherAllow: "claude-opus-4-6",
},
{
// Bedrock request ids are region-prefixed/versioned; the parser
// normalizes them to the catalog key the allowlist holds.
name: "bedrock", catalogID: "bedrock_api", wire: "bedrock",
selectedReq: "us.anthropic.claude-sonnet-4-5-v1:0", selectedAllow: "anthropic.claude-sonnet-4-5",
otherReq: "us.anthropic.claude-opus-4-8-v1:0", otherAllow: "anthropic.claude-opus-4-8",
},
}
vllm, err := harness.StartVLLM(ctx, srv)
require.NoError(t, err, "start mock upstream")
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-pergroup-main"})
require.NoError(t, err, "create main group")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) })
grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-pergroup-other"})
require.NoError(t, err, "create other group")
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) })
ephemeral := false
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
Name: "e2e-pergroup-client",
Type: "reusable",
ExpiresIn: 86400,
UsageLimit: 0,
AutoGroups: []string{grpMain.Id},
Ephemeral: &ephemeral,
})
require.NoError(t, err, "mint setup key")
require.NotEmpty(t, sk.Key, "setup key plaintext")
staticKey := "static-e2e-token"
enabled := true
for i, c := range cases {
req := api.AgentNetworkProviderRequest{
Name: "e2e-pergroup-" + c.name,
ProviderId: c.catalogID,
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
Models: c.models,
}
if i == 0 {
req.BootstrapCluster = ptr(harness.AgentNetworkCluster)
}
prov, perr := srv.CreateProvider(ctx, req)
require.NoError(t, perr, "create provider %s", c.name)
c.providerID = prov.Id
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
gSel := mkAllowGuardrail(t, ctx, "e2e-pergroup-"+c.name+"-sel", c.selectedAllow)
gOth := mkAllowGuardrail(t, ctx, "e2e-pergroup-"+c.name+"-oth", c.otherAllow)
polMain, merr := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-pergroup-" + c.name + "-main",
Enabled: &enabled,
SourceGroups: []string{grpMain.Id},
DestinationProviderIds: []string{prov.Id},
GuardrailIds: &[]string{gSel.Id},
})
require.NoError(t, merr, "create main policy %s", c.name)
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) })
polOther, oerr := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-pergroup-" + c.name + "-other",
Enabled: &enabled,
SourceGroups: []string{grpOther.Id},
DestinationProviderIds: []string{prov.Id},
GuardrailIds: &[]string{gOth.Id},
})
require.NoError(t, oerr, "create other policy %s", c.name)
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) })
}
settings, err := srv.GetSettings(ctx)
require.NoError(t, err, "read settings")
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-pergroup-proxy")
require.NoError(t, err, "mint proxy token")
px, err := harness.StartProxy(ctx, srv, proxyToken)
require.NoError(t, err, "start proxy")
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
cl, err := harness.StartClient(ctx, srv, sk.Key)
require.NoError(t, err, "start client")
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
}
send := func(c *pergroupCase, model string) (int, string) {
var code int
var body string
var cerr error
switch c.wire {
case "vertex":
code, body, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, vertexProject, vertexRegion, model, "Reply with exactly: pong", "")
case "bedrock":
code, body, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, model, "Reply with exactly: pong", "")
default:
code, body, cerr = cl.Chat(ctx, settings.Endpoint, proxyIP, c.wire, model, "Reply with exactly: pong", "")
}
require.NoError(t, cerr, "request must reach the proxy for %s", c.name)
return code, body
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
// grpMain's own model is served. Retry to absorb tunnel/DNS jitter on
// the first call over the freshly warmed tunnel.
var code int
var body string
deadline := time.Now().Add(90 * time.Second)
for time.Now().Before(deadline) {
code, body = send(c, c.selectedReq)
if code == 200 {
break
}
time.Sleep(5 * time.Second)
}
assert.Equal(t, 200, code,
"%s: grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", c.name, body, px.Logs(context.Background()))
// grpOther's model must NOT leak to the grpMain client.
code, body = send(c, c.otherReq)
assert.Equal(t, 403, code,
"%s: another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", c.name, body, px.Logs(context.Background()))
assert.Contains(t, body, "llm_policy.model_blocked",
"%s: denial must be a model-allowlist decision, not routing; body: %s", c.name, body)
})
}
}
// TestGuardrailMultiGroupUser proves the per-policy/group decision for a caller
// that belongs to MULTIPLE groups at once. Two scenarios, one shared stack:
//
// - union across the user's groups: the client is in gUX and gUY, each with
// its own policy+guardrail on provider P1 (gUX->union-a, gUY->union-b). The
// client may use BOTH models (the union of its groups' allowlists) while a
// third, un-allowlisted model is denied.
// - an un-guardrailed group lifts the restriction: the client is in gMP and
// gMQ on provider P2, where gMP restricts to mix-a but gMQ's policy carries
// NO guardrail. Because one applicable policy is unrestricted, the client may
// use a model on no allowlist (mix-z) as well as mix-a.
func TestGuardrailMultiGroupUser(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
defer cancel()
const (
unionA = "mg-union-a"
unionB = "mg-union-b"
unionC = "mg-union-c" // allowlisted by neither group
mixA = "mg-mix-a"
mixZ = "mg-mix-z" // on no allowlist; reachable only via the un-guardrailed policy
)
priced := func(ids ...string) *[]api.AgentNetworkProviderModel {
out := make([]api.AgentNetworkProviderModel, 0, len(ids))
for _, id := range ids {
out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001})
}
return &out
}
vllm, err := harness.StartVLLM(ctx, srv)
require.NoError(t, err, "start mock upstream")
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
mkGroup := func(name string) *api.Group {
g, gerr := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: name})
require.NoError(t, gerr, "create group %s", name)
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), g.Id) })
return g
}
gUX := mkGroup("e2e-mg-union-x")
gUY := mkGroup("e2e-mg-union-y")
gMP := mkGroup("e2e-mg-mix-p")
gMQ := mkGroup("e2e-mg-mix-q")
ephemeral := false
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
Name: "e2e-mg-client",
Type: "reusable",
ExpiresIn: 86400,
UsageLimit: 0,
AutoGroups: []string{gUX.Id, gUY.Id, gMP.Id, gMQ.Id}, // client in all four groups
Ephemeral: &ephemeral,
})
require.NoError(t, err, "mint setup key")
require.NotEmpty(t, sk.Key, "setup key plaintext")
staticKey := "static-e2e-token"
enabled := true
// P1 — union scenario: two restricting policies, one per group.
p1, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
Name: "e2e-mg-union",
ProviderId: "openai_api",
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
Models: priced(unionA, unionB, unionC),
BootstrapCluster: ptr(harness.AgentNetworkCluster),
})
require.NoError(t, err, "create union provider")
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p1.Id) })
polUX, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-mg-union-x",
Enabled: &enabled,
SourceGroups: []string{gUX.Id},
DestinationProviderIds: []string{p1.Id},
GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-union-x", unionA).Id},
})
require.NoError(t, err, "create union policy X")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polUX.Id) })
polUY, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-mg-union-y",
Enabled: &enabled,
SourceGroups: []string{gUY.Id},
DestinationProviderIds: []string{p1.Id},
GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-union-y", unionB).Id},
})
require.NoError(t, err, "create union policy Y")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polUY.Id) })
// P2 — mixed scenario: one restricting policy + one un-guardrailed policy.
p2, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
Name: "e2e-mg-mix",
ProviderId: "openai_api",
UpstreamUrl: vllm.URL,
ApiKey: &staticKey,
Enabled: ptr(true),
Models: priced(mixA, mixZ),
})
require.NoError(t, err, "create mix provider")
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p2.Id) })
polMP, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-mg-mix-p",
Enabled: &enabled,
SourceGroups: []string{gMP.Id},
DestinationProviderIds: []string{p2.Id},
GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-mix-p", mixA).Id},
})
require.NoError(t, err, "create mix policy P")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMP.Id) })
polMQ, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
Name: "e2e-mg-mix-q",
Enabled: &enabled,
SourceGroups: []string{gMQ.Id},
DestinationProviderIds: []string{p2.Id}, // NO guardrail -> unrestricted
})
require.NoError(t, err, "create mix policy Q")
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMQ.Id) })
settings, err := srv.GetSettings(ctx)
require.NoError(t, err, "read settings")
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-mg-proxy")
require.NoError(t, err, "mint proxy token")
px, err := harness.StartProxy(ctx, srv, proxyToken)
require.NoError(t, err, "start proxy")
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
cl, err := harness.StartClient(ctx, srv, sk.Key)
require.NoError(t, err, "start client")
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
}
send := func(model string) (int, string) {
code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "")
require.NoError(t, cerr, "request must reach the proxy")
return code, body
}
sendUntil200 := func(model string) (int, string) {
var code int
var body string
deadline := time.Now().Add(90 * time.Second)
for time.Now().Before(deadline) {
code, body = send(model)
if code == 200 {
break
}
time.Sleep(5 * time.Second)
}
return code, body
}
t.Run("union across the user's groups", func(t *testing.T) {
code, body := sendUntil200(unionA)
assert.Equal(t, 200, code, "model allowed by group X must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
code, body = sendUntil200(unionB)
assert.Equal(t, 200, code, "model allowed by group Y must also be served (union across the user's groups); body: %s", body)
code, body = send(unionC)
assert.Equal(t, 403, code, "a model on neither group's allowlist must be denied; body: %s", body)
assert.Contains(t, body, "llm_policy.model_blocked")
})
t.Run("an un-guardrailed group lifts the restriction", func(t *testing.T) {
code, body := sendUntil200(mixA)
assert.Equal(t, 200, code, "the restricted group's model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
code, body = sendUntil200(mixZ)
assert.Equal(t, 200, code,
"a non-allowlisted model must be served because the user is also in a group whose policy has no guardrail; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
})
}
// mkAllowGuardrail creates a guardrail whose model allowlist is enabled and holds
// exactly the given model, registering cleanup.
func mkAllowGuardrail(t *testing.T, ctx context.Context, name, model string) api.AgentNetworkGuardrail {
t.Helper()
var gr api.AgentNetworkGuardrailRequest
gr.Name = name
gr.Checks.ModelAllowlist.Enabled = true
gr.Checks.ModelAllowlist.Models = []string{model}
g, err := srv.CreateGuardrail(ctx, gr)
require.NoError(t, err, "create guardrail %s", name)
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) })
return g
}

View File

@@ -52,9 +52,7 @@ func catalogModel(pc providerCase) string {
func disallowedModel(pc providerCase) string {
switch pc.kind {
case harness.WireBedrock:
// Same profile prefix as the allowed model so only the model name
// differs; the guardrail must deny it before it reaches AWS.
return strings.SplitN(pc.model, ".", 2)[0] + ".anthropic.claude-opus-4-8"
return "us.anthropic.claude-opus-4-8"
case harness.WireVertex:
return "claude-opus-4-8@20250101"
default:
@@ -74,7 +72,7 @@ func sendModel(ctx context.Context, t *testing.T, cl *harness.Client, endpoint,
case harness.WireVertex:
code, _, err = cl.Vertex(ctx, endpoint, proxyIP, pc.project, pc.region, model, "Reply with exactly: pong", "")
default:
code, _, err = cl.ChatPrefixed(ctx, endpoint, proxyIP, pc.pathPrefix, pc.kind, model, "Reply with exactly: pong", "")
code, _, err = cl.Chat(ctx, endpoint, proxyIP, pc.kind, model, "Reply with exactly: pong", "")
}
require.NoError(t, err, "request must reach the proxy for %s", pc.name)
return code
@@ -166,7 +164,8 @@ func TestModelAllowlistEnforced(t *testing.T) {
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 (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the proxy peer so WaitProxyPeer then observes it connected.
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
require.NoError(t, err, "resolve agent-network endpoint to proxy IP")
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {

View File

@@ -104,7 +104,8 @@ func TestProviderSkipTLSVerification(t *testing.T) {
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 (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the 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 {

View File

@@ -106,7 +106,8 @@ func TestVLLMProvider(t *testing.T) {
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 (DNS error fails) and its first packet wakes the lazy proxy peer, so WaitProxyPeer sees it connected; any HTTP status counts.
// Resolve first: the DNS lookup triggers the lazy-connection warm-up, waking
// the 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 {

View File

@@ -4,7 +4,6 @@ package harness
import (
"context"
"errors"
"fmt"
"io"
"os/exec"
@@ -168,54 +167,22 @@ func (cl *Client) pollStatus(ctx context.Context, timeout time.Duration, want st
return fmt.Errorf("timed out waiting for %q; last status:\n%s", want, last)
}
const (
// curlExitCouldNotResolve is curl's exit code for a DNS resolution failure, distinct from connection-level failures.
curlExitCouldNotResolve = 6
// dnsProbeRetryWindow bounds DNS-failure retries: the synthesized zone lands a beat after management connects, so early NXDOMAIN is propagation; a zone still absent after this window is a real failure.
dnsProbeRetryWindow = 30 * time.Second
dnsProbeRetryInterval = 2 * time.Second
)
// ResolveProxyIP GETs https://<endpoint>/ from the client's netns: any HTTP status proves DNS + tunnel and wakes the lazy proxy peer; only DNS failures retry, within dnsProbeRetryWindow. Returns the connected IP for --resolve pinning.
// ResolveProxyIP resolves the agent-network endpoint to the proxy peer's
// NetBird IP from inside the client (via magic DNS).
func (cl *Client) ResolveProxyIP(ctx context.Context, endpoint string) (string, error) {
args := []string{
"run", "--rm",
"--network", "container:" + cl.container.GetContainerID(),
curlImage,
"-ksS", "-o", "/dev/null",
"--connect-timeout", "30", "--max-time", "60",
"-w", "%{remote_ip}",
"https://" + endpoint + "/",
code, reader, err := cl.container.Exec(ctx, []string{"getent", "hosts", endpoint}, tcexec.Multiplexed())
if err != nil {
return "", err
}
deadline := time.Now().Add(dnsProbeRetryWindow)
for {
cmd := exec.CommandContext(ctx, "docker", args...)
var stdout, stderr strings.Builder
cmd.Stdout = &stdout
cmd.Stderr = &stderr
err := cmd.Run()
if err == nil {
ip := strings.TrimSpace(stdout.String())
if ip == "" {
return "", fmt.Errorf("got an HTTP response from %s but no remote IP", endpoint)
}
return ip, nil
}
var exitErr *exec.ExitError
if !errors.As(err, &exitErr) || exitErr.ExitCode() != curlExitCouldNotResolve {
return "", fmt.Errorf("no HTTP response from %s: %w (%s)", endpoint, err, strings.TrimSpace(stderr.String()))
}
dnsErr := fmt.Errorf("DNS resolution failed for %s: %s", endpoint, strings.TrimSpace(stderr.String()))
if time.Until(deadline) < dnsProbeRetryInterval {
return "", dnsErr
}
select {
case <-ctx.Done():
return "", fmt.Errorf("%w (%w)", dnsErr, ctx.Err())
case <-time.After(dnsProbeRetryInterval):
}
out, _ := io.ReadAll(reader)
if code != 0 {
return "", fmt.Errorf("getent hosts %s exited %d", endpoint, code)
}
fields := strings.Fields(string(out))
if len(fields) == 0 {
return "", fmt.Errorf("no address for %s", endpoint)
}
return fields[0], nil
}
// Wire shapes for Chat.
@@ -239,29 +206,18 @@ const (
// the wire shape: WireChat (OpenAI) or WireMessages (Anthropic). A non-empty
// sessionID is sent as the universal x-session-id header the proxy records.
func (cl *Client) Chat(ctx context.Context, endpoint, proxyIP, kind, model, prompt, sessionID string) (int, string, error) {
return cl.ChatPrefixed(ctx, endpoint, proxyIP, "", kind, model, prompt, sessionID)
}
// ChatPrefixed is Chat with a base-URL path prefix prepended to the wire
// path, mirroring agents whose base URL carries a shape-selecting prefix that
// rides through to the upstream — e.g. Claude Code against a Kimi provider
// sets ANTHROPIC_BASE_URL=https://<endpoint>/anthropic so the proxy forwards
// /anthropic/v1/messages to Moonshot's Anthropic surface while the provider's
// upstream URL stays the bare https://api.moonshot.ai. Empty prefix is plain
// Chat.
func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefix, kind, model, prompt, sessionID string) (int, string, error) {
var path, body string
var headers []string
switch kind {
case WireMessages:
path = "/v1/messages"
headers = []string{"anthropic-version: 2023-06-01"}
body = fmt.Sprintf(`{"model":%q,"max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, model, prompt)
body = fmt.Sprintf(`{"model":%q,"max_tokens":64,"messages":[{"role":"user","content":%q}]}`, model, prompt)
default:
path = "/v1/chat/completions"
body = fmt.Sprintf(`{"model":%q,"messages":[{"role":"user","content":%q}]}`, model, prompt)
}
return cl.post(ctx, endpoint, proxyIP, pathPrefix+path, body, withSessionID(headers, sessionID))
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(headers, sessionID))
}
// Vertex issues an Anthropic-on-Vertex rawPredict POST over the tunnel. Unlike
@@ -271,7 +227,7 @@ func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefi
// is sent as the universal x-session-id header the proxy records.
func (cl *Client) Vertex(ctx context.Context, endpoint, proxyIP, project, region, model, prompt, sessionID string) (int, string, error) {
path := fmt.Sprintf("/v1/projects/%s/locations/%s/publishers/anthropic/models/%s:rawPredict", project, region, model)
body := fmt.Sprintf(`{"anthropic_version":"vertex-2023-10-16","max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, prompt)
body := fmt.Sprintf(`{"anthropic_version":"vertex-2023-10-16","max_tokens":64,"messages":[{"role":"user","content":%q}]}`, prompt)
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(nil, sessionID))
}
@@ -282,7 +238,7 @@ func (cl *Client) Vertex(ctx context.Context, endpoint, proxyIP, project, region
// header the proxy records.
func (cl *Client) Bedrock(ctx context.Context, endpoint, proxyIP, model, prompt, sessionID string) (int, string, error) {
path := "/model/" + model + "/invoke"
body := fmt.Sprintf(`{"anthropic_version":"bedrock-2023-05-31","max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, prompt)
body := fmt.Sprintf(`{"anthropic_version":"bedrock-2023-05-31","max_tokens":64,"messages":[{"role":"user","content":%q}]}`, prompt)
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(nil, sessionID))
}
@@ -341,7 +297,7 @@ func (cl *Client) Terminate(ctx context.Context) error {
return cl.container.Terminate(ctx)
}
// containerLogs reads up to 4 MiB of a container's logs for diagnostics — enough for a whole provider-matrix run.
// containerLogs reads up to 256 KiB of a container's logs for diagnostics.
func containerLogs(ctx context.Context, c testcontainers.Container) string {
if c == nil {
return ""
@@ -351,6 +307,6 @@ func containerLogs(ctx context.Context, c testcontainers.Container) string {
return fmt.Sprintf("<logs error: %v>", err)
}
defer r.Close()
b, _ := io.ReadAll(io.LimitReader(r, 4<<20))
b, _ := io.ReadAll(io.LimitReader(r, 256<<10))
return string(b)
}

View File

@@ -221,29 +221,6 @@ func (c *Combined) CreateProxyTokenCLI(ctx context.Context, name string) (string
return "", fmt.Errorf("token not found in CLI output: %s", string(out))
}
// SnapshotStoreDB copies the management sqlite store (with WAL/SHM sidecars) out of the bind-mounted
// data dir into dstDir and returns the copy's path; reading a copy avoids locking against live writes.
func (c *Combined) SnapshotStoreDB(dstDir string) (string, error) {
src := filepath.Join(c.workDir, "data", "store.db")
if _, err := os.Stat(src); err != nil {
return "", fmt.Errorf("management store not found at %s: %w", src, err)
}
dst := filepath.Join(dstDir, "store.db")
for _, suffix := range []string{"", "-wal", "-shm"} {
data, err := os.ReadFile(src + suffix)
if err != nil {
if os.IsNotExist(err) && suffix != "" {
continue // sidecar only exists in WAL mode
}
return "", fmt.Errorf("read %s: %w", src+suffix, err)
}
if err := os.WriteFile(dst+suffix, data, 0o600); err != nil {
return "", fmt.Errorf("write %s: %w", dst+suffix, err)
}
}
return dst, nil
}
// Logs returns the combined server container logs, for diagnostics.
func (c *Combined) Logs(ctx context.Context) string {
return containerLogs(ctx, c.container)

View File

@@ -38,11 +38,7 @@ type Proxy struct {
// network, registered via the given account proxy token and serving the
// AgentNetworkCluster over a self-signed wildcard cert. It does not wait for
// peer connectivity — callers poll management for the proxy peer.
// StartProxy launches the reverse-proxy container. Optional envOverrides are
// merged into the container environment after the defaults, so callers can set
// or override any NB_PROXY_* var (e.g. NB_PROXY_TUNNEL_CACHE_TTL for tests that
// need a short authorization-cache window).
func StartProxy(ctx context.Context, c *Combined, proxyToken string, envOverrides ...map[string]string) (*Proxy, error) {
func StartProxy(ctx context.Context, c *Combined, proxyToken string) (*Proxy, error) {
root, err := repoRoot()
if err != nil {
return nil, err
@@ -97,12 +93,6 @@ func StartProxy(ctx context.Context, c *Combined, proxyToken string, envOverride
WaitingFor: wait.ForLog("Initial mapping sync complete").WithStartupTimeout(90 * time.Second),
}
for _, ov := range envOverrides {
for k, v := range ov {
req.Env[k] = v
}
}
ctr, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{
ContainerRequest: req,
Started: true,

View File

@@ -234,6 +234,9 @@ init_environment() {
NETBIRD_LICENSE_KEY=$(read_secret "Enter license key (input hidden)")
GHCR_USERNAME="netbirdExtAccess1"
GHCR_TOKEN=$(read_secret "Enter GHCR token (input hidden)")
POSTGRES_USER="netbird"
POSTGRES_DB="netbird"
POSTGRES_PASSWORD=$(rand_secret)
@@ -260,6 +263,10 @@ init_environment() {
install -m 600 /dev/null config.yaml
render_config_yaml >> config.yaml
echo "Logging in to ghcr.io ..."
printf '%s' "$GHCR_TOKEN" | docker login ghcr.io -u "$GHCR_USERNAME" --password-stdin
unset GHCR_TOKEN
echo ""
echo "Pulling images ..."
$DOCKER_COMPOSE_COMMAND pull

View File

@@ -490,6 +490,8 @@ init_migration() {
echo ""
echo "Step 1: Image swap (community → Enterprise). License key required."
NB_LICENSE_KEY=$(read_secret " License key")
GHCR_USERNAME="netbirdExtAccess1"
GHCR_TOKEN=$(read_secret " GHCR token (input hidden)")
# Step 2 — optional
echo ""
@@ -586,6 +588,11 @@ apply_changes() {
fi
} >> "$ENV_FILE"
echo ""
echo "Logging in to ghcr.io ..."
printf '%s' "$GHCR_TOKEN" | docker login ghcr.io -u "$GHCR_USERNAME" --password-stdin
unset GHCR_TOKEN
echo ""
echo "Pulling enterprise images ..."
$DOCKER_COMPOSE_COMMAND pull

View File

@@ -18,26 +18,21 @@ import (
// contract between the proxy and management; management flattens them into
// queryable columns. Keep in sync with the proxy side.
const (
metaKeyProvider = "llm.provider"
metaKeyModel = "llm.model"
metaKeyResolvedProviderID = "llm.resolved_provider_id"
metaKeySelectedPolicyID = "llm.selected_policy_id"
metaKeyPolicyDecision = "llm_policy.decision"
metaKeyPolicyReason = "llm_policy.reason"
metaKeyInputTokens = "llm.input_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyOutputTokens = "llm.output_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyTotalTokens = "llm.total_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCachedInputTokens = "llm.cached_input_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCacheCreationTokens = "llm.cache_creation_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCostUSDInput = "cost.usd_input"
metaKeyCostUSDCachedInput = "cost.usd_cached_input"
metaKeyCostUSDCacheCreate = "cost.usd_cache_creation"
metaKeyCostUSDOutput = "cost.usd_output"
metaKeyStream = "llm.stream"
metaKeySessionID = "llm.session_id"
metaKeyAuthorisingGroups = "llm.authorising_groups"
metaKeyRequestPrompt = "llm.request_prompt"
metaKeyResponseCompletion = "llm.response_completion"
metaKeyProvider = "llm.provider"
metaKeyModel = "llm.model"
metaKeyResolvedProviderID = "llm.resolved_provider_id"
metaKeySelectedPolicyID = "llm.selected_policy_id"
metaKeyPolicyDecision = "llm_policy.decision"
metaKeyPolicyReason = "llm_policy.reason"
metaKeyInputTokens = "llm.input_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyOutputTokens = "llm.output_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyTotalTokens = "llm.total_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCostUSDTotal = "cost.usd_total"
metaKeyStream = "llm.stream"
metaKeySessionID = "llm.session_id"
metaKeyAuthorisingGroups = "llm.authorising_groups"
metaKeyRequestPrompt = "llm.request_prompt"
metaKeyResponseCompletion = "llm.response_completion"
)
// IngestAccessLog flattens the metadata-bearing reverse-proxy access-log entry
@@ -113,25 +108,20 @@ func flattenAccessLog(e *accesslogs.AccessLogEntry) (*types.AgentNetworkAccessLo
BytesUpload: e.BytesUpload,
BytesDownload: e.BytesDownload,
Provider: meta[metaKeyProvider],
Model: meta[metaKeyModel],
SessionID: meta[metaKeySessionID],
ResolvedProviderID: meta[metaKeyResolvedProviderID],
SelectedPolicyID: meta[metaKeySelectedPolicyID],
Decision: meta[metaKeyPolicyDecision],
DenyReason: meta[metaKeyPolicyReason],
InputTokens: parseMetaInt(meta, metaKeyInputTokens),
OutputTokens: parseMetaInt(meta, metaKeyOutputTokens),
TotalTokens: parseMetaInt(meta, metaKeyTotalTokens),
CachedInputTokens: parseMetaInt(meta, metaKeyCachedInputTokens),
CacheCreationTokens: parseMetaInt(meta, metaKeyCacheCreationTokens),
InputCostUSD: parseMetaFloat(meta, metaKeyCostUSDInput),
CachedInputCostUSD: parseMetaFloat(meta, metaKeyCostUSDCachedInput),
CacheCreationCostUSD: parseMetaFloat(meta, metaKeyCostUSDCacheCreate),
OutputCostUSD: parseMetaFloat(meta, metaKeyCostUSDOutput),
Stream: parseMetaBool(meta, metaKeyStream),
RequestPrompt: meta[metaKeyRequestPrompt],
ResponseCompletion: meta[metaKeyResponseCompletion],
Provider: meta[metaKeyProvider],
Model: meta[metaKeyModel],
SessionID: meta[metaKeySessionID],
ResolvedProviderID: meta[metaKeyResolvedProviderID],
SelectedPolicyID: meta[metaKeySelectedPolicyID],
Decision: meta[metaKeyPolicyDecision],
DenyReason: meta[metaKeyPolicyReason],
InputTokens: parseMetaInt(meta, metaKeyInputTokens),
OutputTokens: parseMetaInt(meta, metaKeyOutputTokens),
TotalTokens: parseMetaInt(meta, metaKeyTotalTokens),
CostUSD: parseMetaFloat(meta, metaKeyCostUSDTotal),
Stream: parseMetaBool(meta, metaKeyStream),
RequestPrompt: meta[metaKeyRequestPrompt],
ResponseCompletion: meta[metaKeyResponseCompletion],
}
var groups []types.AgentNetworkAccessLogGroup
@@ -150,23 +140,18 @@ func flattenAccessLog(e *accesslogs.AccessLogEntry) (*types.AgentNetworkAccessLo
// log's ID so the two correlate.
func usageFromFlattenedLog(e *types.AgentNetworkAccessLog, groups []types.AgentNetworkAccessLogGroup) (*types.AgentNetworkUsage, []types.AgentNetworkUsageGroup) {
usage := &types.AgentNetworkUsage{
ID: e.ID,
AccountID: e.AccountID,
Timestamp: e.Timestamp,
UserID: e.UserID,
ResolvedProviderID: e.ResolvedProviderID,
Provider: e.Provider,
Model: e.Model,
SessionID: e.SessionID,
InputTokens: e.InputTokens,
OutputTokens: e.OutputTokens,
TotalTokens: e.TotalTokens,
CachedInputTokens: e.CachedInputTokens,
CacheCreationTokens: e.CacheCreationTokens,
InputCostUSD: e.InputCostUSD,
CachedInputCostUSD: e.CachedInputCostUSD,
CacheCreationCostUSD: e.CacheCreationCostUSD,
OutputCostUSD: e.OutputCostUSD,
ID: e.ID,
AccountID: e.AccountID,
Timestamp: e.Timestamp,
UserID: e.UserID,
ResolvedProviderID: e.ResolvedProviderID,
Provider: e.Provider,
Model: e.Model,
SessionID: e.SessionID,
InputTokens: e.InputTokens,
OutputTokens: e.OutputTokens,
TotalTokens: e.TotalTokens,
CostUSD: e.CostUSD,
}
usageGroups := make([]types.AgentNetworkUsageGroup, 0, len(groups))

View File

@@ -28,22 +28,17 @@ func newIngestTestEntry() *accesslogs.AccessLogEntry {
UserId: "user-1",
AgentNetwork: true,
Metadata: map[string]string{
metaKeyProvider: "openai",
metaKeyModel: "gpt-5.4",
metaKeyResolvedProviderID: "prov-1",
metaKeySessionID: "sess-1",
metaKeyInputTokens: "100",
metaKeyOutputTokens: "50",
metaKeyTotalTokens: "1174",
metaKeyCachedInputTokens: "256",
metaKeyCacheCreationTokens: "768",
metaKeyCostUSDInput: "0.0071",
metaKeyCostUSDCachedInput: "0.0009",
metaKeyCostUSDCacheCreate: "0.0020",
metaKeyCostUSDOutput: "0.0023",
metaKeyStream: "true",
metaKeyRequestPrompt: "hello",
metaKeyResponseCompletion: "world",
metaKeyProvider: "openai",
metaKeyModel: "gpt-5.4",
metaKeyResolvedProviderID: "prov-1",
metaKeySessionID: "sess-1",
metaKeyInputTokens: "100",
metaKeyOutputTokens: "50",
metaKeyTotalTokens: "150",
metaKeyCostUSDTotal: "0.0123",
metaKeyStream: "true",
metaKeyRequestPrompt: "hello",
metaKeyResponseCompletion: "world",
// repeated id must be de-duplicated before the group rows insert.
metaKeyAuthorisingGroups: "grp-eng,grp-eng,grp-ops",
},
@@ -70,19 +65,7 @@ func TestIngestAccessLog_RealStore_LogCollectionOff(t *testing.T) {
require.Len(t, usage, 1, "usage row must be written even with log collection off")
assert.Equal(t, int64(100), usage[0].InputTokens, "input tokens must round-trip from metadata")
assert.Equal(t, int64(50), usage[0].OutputTokens, "output tokens must round-trip from metadata")
assert.Equal(t, int64(256), usage[0].CachedInputTokens, "cache-read tokens must round-trip from metadata")
assert.Equal(t, int64(768), usage[0].CacheCreationTokens, "cache-write tokens must round-trip from metadata")
// The per-bucket breakdown is the only cost state stored, and must survive
// the write/read cycle as real columns — usage rows are the only cost
// record for accounts with log collection off, so a dropped column here
// loses the split permanently.
assert.InDelta(t, 0.0071, usage[0].InputCostUSD, 1e-9, "input cost must round-trip from metadata")
assert.InDelta(t, 0.0009, usage[0].CachedInputCostUSD, 1e-9, "cache-read cost must round-trip from metadata")
assert.InDelta(t, 0.0020, usage[0].CacheCreationCostUSD, 1e-9, "cache-write cost must round-trip from metadata")
assert.InDelta(t, 0.0023, usage[0].OutputCostUSD, 1e-9, "output cost must round-trip from metadata")
// Aggregates are derived from the stored columns, never stored themselves.
assert.InDelta(t, 0.0123, usage[0].TotalCostUSD(), 1e-9, "total is derived from the stored buckets")
assert.InDelta(t, 0.0029, usage[0].CacheCostUSD(), 1e-9, "cache cost is derived from the two cache buckets")
assert.InDelta(t, 0.0123, usage[0].CostUSD, 1e-9, "cost must round-trip from metadata")
logs, _, err := s.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, testAccountID, types.AgentNetworkAccessLogFilter{})
require.NoError(t, err)
@@ -113,14 +96,6 @@ func TestIngestAccessLog_RealStore_LogCollectionOn(t *testing.T) {
require.Equal(t, int64(1), total, "exactly one access-log row expected")
require.Len(t, logs, 1, "full access-log row must be written when log collection is on")
assert.Equal(t, "gpt-5.4", logs[0].Model, "model must flatten from metadata")
assert.Equal(t, int64(256), logs[0].CachedInputTokens, "cache-read tokens must flatten from metadata")
assert.Equal(t, int64(768), logs[0].CacheCreationTokens, "cache-write tokens must flatten from metadata")
assert.InDelta(t, 0.0029, logs[0].CacheCostUSD(), 1e-9, "cache cost is derived from the two cache buckets")
assert.InDelta(t, 0.0123, logs[0].TotalCostUSD(), 1e-9, "total is derived from the stored buckets")
assert.InDelta(t, 0.0071, logs[0].InputCostUSD, 1e-9, "input cost must flatten from metadata")
assert.InDelta(t, 0.0009, logs[0].CachedInputCostUSD, 1e-9, "cache-read cost must flatten from metadata")
assert.InDelta(t, 0.0020, logs[0].CacheCreationCostUSD, 1e-9, "cache-write cost must flatten from metadata")
assert.InDelta(t, 0.0023, logs[0].OutputCostUSD, 1e-9, "output cost must flatten from metadata")
assert.Equal(t, "hello", logs[0].RequestPrompt, "prompt must be retained when log collection is on")
assert.Equal(t, "world", logs[0].ResponseCompletion, "completion must be retained when log collection is on")
assert.True(t, logs[0].Stream, "stream flag must flatten from metadata")

View File

@@ -38,7 +38,7 @@ func accessLogRow(id, sessionID string, ts time.Time, opts ...func(*types.AgentN
InputTokens: 100,
OutputTokens: 50,
TotalTokens: 150,
InputCostUSD: 0.01,
CostUSD: 0.01,
}
for _, o := range opts {
o(e)
@@ -74,7 +74,7 @@ func withTokens(in, out, total int64, cost float64) func(*types.AgentNetworkAcce
e.InputTokens = in
e.OutputTokens = out
e.TotalTokens = total
e.InputCostUSD = cost
e.CostUSD = cost
}
}
@@ -155,7 +155,7 @@ func TestAccessLogSessions_FoldAndAggregate(t *testing.T) {
assert.Equal(t, int64(310), a.InputTokens, "input tokens summed")
assert.Equal(t, int64(135), a.OutputTokens, "output tokens summed")
assert.Equal(t, int64(445), a.TotalTokens, "total tokens summed")
assert.InDelta(t, 0.031, a.TotalCostUSD(), 1e-9, "cost summed")
assert.InDelta(t, 0.031, a.CostUSD, 1e-9, "cost summed")
assert.Equal(t, "deny", a.Decision, "any deny makes the session a deny")
assert.ElementsMatch(t, []string{"openai", "anthropic"}, a.Providers, "distinct providers")
assert.ElementsMatch(t, []string{"gpt-5.4", "claude-haiku-4-5"}, a.Models, "distinct models")

View File

@@ -420,47 +420,6 @@ var providers = []Provider{
{ID: "mistral-embed", Label: "Mistral Embed", InputPer1k: 0.0001, OutputPer1k: 0, ContextWindow: 8192},
},
},
{
ID: "kimi_api",
Kind: KindProvider,
Name: "Kimi (Moonshot AI) API",
Description: "Kimi K3 / K2 models via the Moonshot AI platform",
DefaultHost: "api.moonshot.ai",
AuthHeaderName: "Authorization",
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#1A1A2E",
// ParserID empty on purpose: Moonshot serves two body shapes on
// the same host and key, and the proxy's URL sniffer dispatches
// both (same pattern as Bifrost). /v1/chat/completions matches
// OpenAIParser; the Anthropic-compatible endpoint the official
// Claude Code guide uses (/anthropic/v1/messages) contains
// "/v1/messages" and matches AnthropicParser. Pinning "openai"
// here would misparse the Claude Code path — the primary way
// teams consume Kimi for coding today. Both endpoints accept the
// same Moonshot key via Authorization: Bearer (Claude Code's
// ANTHROPIC_AUTH_TOKEN rides that header too).
//
// api.moonshot.ai is the international platform; mainland-China
// accounts live on api.moonshot.cn with separate billing —
// operators there override the host on the provider record. The
// kimi.com subscription coding endpoint (api.kimi.com/coding,
// model id "k3") is account-bound seat licensing rather than a
// meterable platform key, so it's deliberately not the default.
ParserID: "",
// 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
// accounts — K2-era ids (kimi-k2-thinking) and even the kimi-latest
// alias return resource_not_found_error, verified live 2026-07-21 —
// so it's the only catalog entry. Grandfathered accounts with K2
// access can still type those ids on the provider's model rows.
// 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},
},
},
{
ID: "litellm_proxy",
Kind: KindGateway,

View File

@@ -86,17 +86,12 @@ type Manager interface {
// PolicySelectionInput is the per-request selection envelope. The
// proxy populates it from CapturedData (account, user, groups) plus
// the provider llm_router resolved and the model it extracted.
// the provider llm_router resolved.
type PolicySelectionInput struct {
AccountID string
UserID string
GroupIDs []string
ProviderID string
// Model is the already-normalised upstream model id the proxy extracted
// (parser strips Bedrock region/version, Vertex @version), so a
// case-insensitive compare suffices. Empty = undetermined → not permitted
// (fail closed).
Model string
}
// PolicySelectionResult names the policy that "pays" for this request

View File

@@ -5,7 +5,6 @@ import (
"fmt"
"math"
"sort"
"strings"
"time"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
@@ -36,10 +35,6 @@ const (
denyCodeAccountTokenCapExceeded = "llm_account.token_cap_exceeded"
//nolint:gosec // account deny code label, not a credential
denyCodeAccountBudgetCapExceeded = "llm_account.budget_cap_exceeded"
// denyCodeModelBlocked is returned when policies govern the request's
// (provider, caller-groups) but none permits the model. Matches the proxy
// guardrail's code so both layers surface the same label.
denyCodeModelBlocked = "llm_policy.model_blocked"
)
// consumptionCache holds the consumption counters prefetched for one
@@ -164,25 +159,6 @@ func (m *managerImpl) SelectPolicyForRequest(ctx context.Context, in PolicySelec
}
candidates := filterApplicablePolicies(policies, in)
// Model-allowlist gate scoped to the matched policies: keep candidates whose
// guardrails permit the model (none enabled = unrestricted), deny when
// policies apply but none permits it. Skip the load when none has a guardrail.
if len(candidates) > 0 && anyPolicyHasGuardrails(candidates) {
guardrailsByID, gErr := m.loadGuardrailsByID(ctx, in.AccountID)
if gErr != nil {
return nil, gErr
}
permitted := filterModelPermittedPolicies(candidates, guardrailsByID, in.Model)
if len(permitted) == 0 {
return &PolicySelectionResult{
Allow: false,
DenyCode: denyCodeModelBlocked,
DenyReason: modelBlockedReason(in.Model),
}, nil
}
candidates = permitted
}
// Prefetch every consumption counter the ceiling + candidate policies will
// read, in a single store round-trip, then score against the cache.
cache, err := m.prefetchConsumption(ctx, in, rules, candidates, now)
@@ -274,90 +250,6 @@ func filterApplicablePolicies(policies []*types.Policy, in PolicySelectionInput)
return out
}
// anyPolicyHasGuardrails reports whether any policy references at least one
// guardrail, so the selector can skip loading guardrails when none do.
func anyPolicyHasGuardrails(policies []*types.Policy) bool {
for _, p := range policies {
if p != nil && len(p.GuardrailIDs) > 0 {
return true
}
}
return false
}
// loadGuardrailsByID loads the account's guardrails indexed by ID. Used by the
// model-allowlist gate to resolve each candidate policy's attached guardrails.
func (m *managerImpl) loadGuardrailsByID(ctx context.Context, accountID string) (map[string]*types.Guardrail, error) {
guardrails, err := m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID)
if err != nil {
return nil, fmt.Errorf("list account guardrails: %w", err)
}
byID := make(map[string]*types.Guardrail, len(guardrails))
for _, g := range guardrails {
if g != nil {
byID[g.ID] = g
}
}
return byID, nil
}
// filterModelPermittedPolicies returns the subset of policies whose guardrails
// permit the model. Order is preserved so downstream scoring is unaffected.
func filterModelPermittedPolicies(policies []*types.Policy, byID map[string]*types.Guardrail, model string) []*types.Policy {
out := make([]*types.Policy, 0, len(policies))
for _, p := range policies {
if policyPermitsModel(p, byID, model) {
out = append(out, p)
}
}
return out
}
// policyPermitsModel reports whether a policy permits the model. No
// allowlist-enabled guardrail = unrestricted (permits any, incl. empty);
// otherwise the model must be in the union of its allowlists, so an
// empty/undetermined model fails closed.
func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model string) bool {
if p == nil {
return false
}
wanted := normaliseModelID(model)
restricted := false
for _, gID := range p.GuardrailIDs {
g, ok := byID[gID]
if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled {
continue
}
restricted = true
if wanted == "" {
continue
}
for _, allowed := range g.Checks.ModelAllowlist.Models {
if normaliseModelID(allowed) == wanted {
return true
}
}
}
return !restricted
}
// normaliseModelID lowercases and trims a model identifier so the allowlist
// compare is case-insensitive and trim-tolerant. Mirrors the proxy guardrail's
// normaliseModel so both layers agree on what "same model" means.
func normaliseModelID(model string) string {
return strings.ToLower(strings.TrimSpace(model))
}
// modelBlockedReason builds the human-readable deny reason for a model-allowlist
// rejection. The model is quoted when known; an undetermined model is reported
// as such so the access log distinguishes "wrong model" from "no model".
func modelBlockedReason(model string) string {
if normaliseModelID(model) == "" {
return "request model could not be determined for the policy allowlist"
}
return fmt.Sprintf("model %q is not permitted by any applicable policy allowlist", model)
}
// candidate is the per-policy intermediate the selector ranks. A
// policy that's been exhausted on any enabled cap never makes it
// into this slice; the selector's deny envelope carries the latest

View File

@@ -1,329 +0,0 @@
package agentnetwork
import (
"context"
"errors"
"testing"
"time"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/store"
)
// guardedPolicy builds an enabled, uncapped policy that authorises sourceGroups
// to reach providerID under the given guardrails. Uncapped keeps the selector's
// headroom scoring trivial so these tests isolate the model-allowlist gate.
func guardedPolicy(id, account string, sourceGroups []string, providerID string, guardrailIDs ...string) *types.Policy {
return &types.Policy{
ID: id,
AccountID: account,
Enabled: true,
SourceGroups: sourceGroups,
DestinationProviderIDs: []string{providerID},
GuardrailIDs: guardrailIDs,
CreatedAt: time.Now().UTC(),
}
}
// allowlistGuardrail builds a guardrail whose model allowlist is enabled and
// carries the given models.
func allowlistGuardrail(id, account string, models ...string) *types.Guardrail {
return &types.Guardrail{
ID: id,
AccountID: account,
Checks: types.GuardrailChecks{
ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true, Models: models},
},
}
}
func expectPolicies(mockStore *store.MockStore, account string, policies ...*types.Policy) {
mockStore.EXPECT().
GetAccountAgentNetworkPolicies(gomock.Any(), gomock.Any(), account).
Return(policies, nil)
}
func expectGuardrails(mockStore *store.MockStore, account string, guardrails ...*types.Guardrail) {
mockStore.EXPECT().
GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), account).
Return(guardrails, nil)
}
// TestSelectPolicy_ModelBlockedByAllowlist proves the authoritative allowlist
// decision: a policy authorises the (provider, group) but restricts the model,
// and the requested model isn't on the list, so the request is denied.
func TestSelectPolicy_ModelBlockedByAllowlist(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o"))
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "claude-opus-4",
})
require.NoError(t, err)
assert.False(t, res.Allow, "a model outside the only applicable policy's allowlist must be denied")
assert.Equal(t, denyCodeModelBlocked, res.DenyCode, "deny code must be model_blocked")
assert.NotEmpty(t, res.DenyReason, "deny reason must be populated")
}
// TestSelectPolicy_ModelAllowedByAllowlist is the allow counterpart: the model
// is on the applicable policy's allowlist, so selection proceeds normally.
func TestSelectPolicy_ModelAllowedByAllowlist(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o", "claude-opus-4"))
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
UserID: "user-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "claude-opus-4",
})
require.NoError(t, err)
assert.True(t, res.Allow, "a model on the applicable policy's allowlist must be allowed")
assert.Equal(t, "pol-A", res.SelectedPolicyID)
}
// TestSelectPolicy_CaseInsensitiveModelMatch proves the compare tolerates case
// and surrounding whitespace, matching the proxy guardrail's normalisation.
func TestSelectPolicy_CaseInsensitiveModelMatch(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", " GPT-4o "))
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "gpt-4o",
})
require.NoError(t, err)
assert.True(t, res.Allow, "case/whitespace variants must match the allowlist entry")
}
// TestSelectPolicy_UnguardedPolicyIsUnrestricted is the false-deny fix: when two
// policies authorise the same (provider, group) and one has no guardrail, that
// policy makes the request unrestricted — not caught by the other's allowlist.
func TestSelectPolicy_UnguardedPolicyIsUnrestricted(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
restricted := guardedPolicy("pol-restricted", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
open := guardedPolicy("pol-open", "acc-1", []string{"grp-eng"}, "prov-1") // no guardrail
expectPolicies(mockStore, "acc-1", restricted, open)
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o"))
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "claude-opus-4",
})
require.NoError(t, err)
assert.True(t, res.Allow, "an un-guardrailed policy for the same (provider, group) must leave the request unrestricted")
assert.Equal(t, "pol-open", res.SelectedPolicyID, "the unrestricted policy must be the one that pays")
}
// TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups is the false-allow fix: a
// model allowlisted only for grp-b must not be usable by a grp-a caller. The
// selector considers only policies applicable to the caller's groups.
func TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
polA := guardedPolicy("pol-a", "acc-1", []string{"grp-a"}, "prov-1", "g-a")
polB := guardedPolicy("pol-b", "acc-1", []string{"grp-b"}, "prov-1", "g-b")
expectPolicies(mockStore, "acc-1", polA, polB)
expectGuardrails(mockStore, "acc-1",
allowlistGuardrail("g-a", "acc-1", "gpt-4o"),
allowlistGuardrail("g-b", "acc-1", "claude-opus-4"),
)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-a"},
ProviderID: "prov-1",
Model: "claude-opus-4", // only allowed for grp-b
})
require.NoError(t, err)
assert.False(t, res.Allow, "grp-b's allowlisted model must not leak to a grp-a caller")
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
}
// TestSelectPolicy_UndeterminedModelFailsClosed proves the fail-closed contract
// mirrors the proxy: with a restricted applicable policy and an empty model
// (e.g. a path-routed shape the parser couldn't map), the request is denied.
func TestSelectPolicy_UndeterminedModelFailsClosed(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o"))
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "", // undetermined
})
require.NoError(t, err)
assert.False(t, res.Allow, "an undetermined model must fail closed against a restricted policy")
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
}
// TestSelectPolicy_DisabledAllowlistDoesNotRestrict proves a guardrail whose
// model allowlist is disabled imposes no model restriction, even though the
// policy references it.
func TestSelectPolicy_DisabledAllowlistDoesNotRestrict(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
disabled := &types.Guardrail{
ID: "g-1",
AccountID: "acc-1",
Checks: types.GuardrailChecks{
ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}},
},
}
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1", disabled)
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "anything-goes",
})
require.NoError(t, err)
assert.True(t, res.Allow, "a disabled allowlist must not restrict the model")
assert.Equal(t, "pol-A", res.SelectedPolicyID)
}
// TestSelectPolicy_UnionAcrossPolicyGuardrails proves a policy with multiple
// allowlist guardrails permits the union of their models (not just the first).
func TestSelectPolicy_UnionAcrossPolicyGuardrails(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1", "g-2")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1",
allowlistGuardrail("g-1", "acc-1", "gpt-4o"),
allowlistGuardrail("g-2", "acc-1", "claude-opus-4"),
)
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "claude-opus-4", // only in the second guardrail's list
})
require.NoError(t, err)
assert.True(t, res.Allow, "a model in any of the policy's allowlist guardrails must be permitted")
assert.Equal(t, "pol-A", res.SelectedPolicyID)
}
// TestSelectPolicy_GuardrailLookupErrorPropagates proves a store failure while
// resolving the candidate policies' guardrails surfaces as an error, not a
// silent allow/deny.
func TestSelectPolicy_GuardrailLookupErrorPropagates(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
expectPolicies(mockStore, "acc-1", policy)
mockStore.EXPECT().
GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), "acc-1").
Return(nil, errors.New("store unavailable"))
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "gpt-4o",
})
require.Error(t, err, "a guardrail-lookup failure must surface as an error")
assert.Nil(t, res)
}
// TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted proves a
// policy referencing a guardrail ID absent from the account's set (a stale
// reference) imposes no model restriction — same as no guardrail.
func TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-missing")
expectPolicies(mockStore, "acc-1", policy)
expectGuardrails(mockStore, "acc-1")
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "anything-goes",
})
require.NoError(t, err)
assert.True(t, res.Allow, "an orphaned guardrail reference must not restrict the model")
assert.Equal(t, "pol-A", res.SelectedPolicyID)
}
// TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter proves the model
// gate narrows candidates before cap scoring: the permitting policy is selected
// even though the blocked one has a larger, more attractive cap.
func TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter(t *testing.T) {
ctrl := gomock.NewController(t)
mgr, mockStore := newSelectorMgr(t, ctrl)
polBig := guardedPolicy("pol-big", "acc-1", []string{"grp-eng"}, "prov-1", "g-restrict")
polBig.Limits = types.PolicyLimits{
TokenLimit: types.PolicyTokenLimit{Enabled: true, GroupCap: 1_000_000, WindowSeconds: 3600},
}
polSmall := guardedPolicy("pol-small", "acc-1", []string{"grp-eng"}, "prov-1", "g-permit")
polSmall.Limits = types.PolicyLimits{
TokenLimit: types.PolicyTokenLimit{Enabled: true, GroupCap: 100, WindowSeconds: 3600},
}
expectPolicies(mockStore, "acc-1", polBig, polSmall)
expectGuardrails(mockStore, "acc-1",
allowlistGuardrail("g-restrict", "acc-1", "gpt-4o"),
allowlistGuardrail("g-permit", "acc-1", "claude-opus-4"),
)
expectConsumptionBatch(mockStore, nil)
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
AccountID: "acc-1",
GroupIDs: []string{"grp-eng"},
ProviderID: "prov-1",
Model: "claude-opus-4", // only pol-small's guardrail permits this
})
require.NoError(t, err)
assert.True(t, res.Allow)
assert.Equal(t, "pol-small", res.SelectedPolicyID,
"the model filter must exclude pol-big before cap scoring")
}

View File

@@ -235,12 +235,7 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
applyAccountCollectionControls(&mergedGuardrails, settings)
// The proxy guardrail is a per-provider fail-closed backstop; the
// authoritative per-policy/group decision is management's
// SelectPolicyForRequest. A provider lands in this map only when every
// authorising policy restricts models.
providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID)
guardrailJSON, err := marshalGuardrailConfig(providerAllowlists, mergedGuardrails.PromptCapture)
guardrailJSON, err := marshalGuardrailConfig(mergedGuardrails)
if err != nil {
return nil, err
}
@@ -785,12 +780,10 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt
// guardrailConfig is the JSON shape the proxy-side llm_guardrail
// middleware expects. Mirrors the proxy registration documented in
// the management→proxy contract. provider_allowlists is keyed by the
// resolved provider id llm_router stamps; a provider absent from the map is
// unrestricted at the proxy layer.
// the management→proxy contract.
type guardrailConfig struct {
ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"`
PromptCapture guardrailPromptCapture `json:"prompt_capture"`
ModelAllowlist []string `json:"model_allowlist,omitempty"`
PromptCapture guardrailPromptCapture `json:"prompt_capture"`
}
type guardrailPromptCapture struct {
@@ -835,10 +828,13 @@ func applyAccountCollectionControls(merged *MergedGuardrails, settings *types.Se
merged.PromptCapture.RedactPii = settings.RedactPii || merged.PromptCapture.RedactPii
}
func marshalGuardrailConfig(providerAllowlists map[string][]string, capture MergedPromptCapture) ([]byte, error) {
func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) {
cfg := guardrailConfig{
ProviderAllowlists: providerAllowlists,
PromptCapture: guardrailPromptCapture(capture),
ModelAllowlist: merged.ModelAllowlist,
PromptCapture: guardrailPromptCapture{
Enabled: merged.PromptCapture.Enabled,
RedactPii: merged.PromptCapture.RedactPii,
},
}
out, err := json.Marshal(cfg)
if err != nil {
@@ -847,74 +843,6 @@ func marshalGuardrailConfig(providerAllowlists map[string][]string, capture Merg
return out, nil
}
// buildProviderAllowlists returns the proxy's per-provider backstop: a provider
// is included only when every authorising policy restricts models (their union);
// if any leaves it unrestricted it is omitted, so management decides per group.
func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Guardrail) map[string][]string {
type providerAcc struct {
models map[string]struct{}
anyUnrestricted bool
}
accs := make(map[string]*providerAcc)
for _, p := range policies {
if p == nil {
continue
}
restricted, models := policyModelAllowlist(p, byID)
for _, providerID := range p.DestinationProviderIDs {
if providerID == "" {
continue
}
acc, ok := accs[providerID]
if !ok {
acc = &providerAcc{models: make(map[string]struct{})}
accs[providerID] = acc
}
if !restricted {
acc.anyUnrestricted = true
continue
}
for _, m := range models {
acc.models[m] = struct{}{}
}
}
}
out := make(map[string][]string, len(accs))
for providerID, acc := range accs {
if acc.anyUnrestricted {
continue
}
models := make([]string, 0, len(acc.models))
for m := range acc.models {
models = append(models, m)
}
sort.Strings(models)
out[providerID] = models
}
return out
}
// policyModelAllowlist reports whether a policy restricts models (has an
// allowlist-enabled guardrail) and the union of allowed models. Models are
// verbatim; the proxy factory lowercases/trims them at decode time.
func policyModelAllowlist(p *types.Policy, byID map[string]*types.Guardrail) (bool, []string) {
restricted := false
var models []string
for _, gID := range p.GuardrailIDs {
g, ok := byID[gID]
if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled {
continue
}
restricted = true
for _, m := range g.Checks.ModelAllowlist.Models {
if m != "" {
models = append(models, m)
}
}
}
return restricted, models
}
// buildAccountService composes the per-account gateway Service. The
// target carries the noop placeholder URL — the router middleware
// rewrites every request to the matched provider's upstream before the
@@ -1058,11 +986,38 @@ func unionSourceGroups(policies []*types.Policy) []string {
return out
}
// MergedGuardrails is the synthesiser's fold target. Only prompt capture is
// merged here — the model allowlist is emitted per-provider, and
// token/budget/retention moved onto Policy.Limits and account Settings.
// MergedGuardrails is the JSON shape passed to the proxy via the
// guardrail middleware's config_json. Mirrors the proxy-side
// expectations and is intentionally distinct from
// types.GuardrailChecks so we can evolve either side independently.
type MergedGuardrails struct {
PromptCapture MergedPromptCapture
ModelAllowlist []string `json:"model_allowlist,omitempty"`
TokenLimits MergedTokenLimits `json:"token_limits"`
Budget MergedBudget `json:"budget"`
PromptCapture MergedPromptCapture `json:"prompt_capture"`
Retention MergedRetention `json:"retention"`
}
type MergedTokenLimits struct {
Hourly *MergedTokenWindow `json:"hourly,omitempty"`
Daily *MergedTokenWindow `json:"daily,omitempty"`
Monthly *MergedTokenWindow `json:"monthly,omitempty"`
}
type MergedTokenWindow struct {
MaxInputTokens int `json:"max_input_tokens,omitempty"`
MaxOutputTokens int `json:"max_output_tokens,omitempty"`
}
type MergedBudget struct {
Hourly *MergedBudgetWindow `json:"hourly,omitempty"`
Daily *MergedBudgetWindow `json:"daily,omitempty"`
Monthly *MergedBudgetWindow `json:"monthly,omitempty"`
}
type MergedBudgetWindow struct {
SoftCapUSD float64 `json:"soft_cap_usd,omitempty"`
HardCapUSD float64 `json:"hard_cap_usd,omitempty"`
}
type MergedPromptCapture struct {
@@ -1070,31 +1025,64 @@ type MergedPromptCapture struct {
RedactPii bool `json:"redact_pii"`
}
// mergeGuardrails folds the referencing policies' guardrails into the
// prompt-capture decision only. The model allowlist is enforced per-policy/group
// in management and shipped per-provider; token/budget/retention live off
// guardrails now.
type MergedRetention struct {
Enabled bool `json:"enabled"`
Days int `json:"days"`
}
// mergeGuardrails computes the effective guardrail spec applied at the
// proxy, given the referencing policies and the account's guardrail
// catalogue. Policy enabled-ness is the caller's responsibility — only
// enabled policies should be passed in.
//
// Merge rule — prompt capture: enabled if any policy enables it; redact_pii
// sticks if any enabling policy turns it on.
// Merge rules:
// - Model allowlist: union of allowlists across policies that enable it.
// - Token / Budget: most-restrictive (min of non-zero caps) per window.
// - Prompt capture: enabled if any policy enables it; redact_pii sticks
// if any enabling policy turns it on.
// - Retention: enabled if any enables it; smallest non-zero days wins.
func mergeGuardrails(policies []*types.Policy, byID map[string]*types.Guardrail) MergedGuardrails {
merged := MergedGuardrails{}
allowlist := make(map[string]struct{})
allowlistEnabled := false
for _, policy := range policies {
for _, gID := range policy.GuardrailIDs {
g, ok := byID[gID]
if !ok || g == nil {
continue
}
mergeGuardrail(g, &merged)
mergeGuardrail(g, &merged, allowlist, &allowlistEnabled)
}
}
if allowlistEnabled {
merged.ModelAllowlist = make([]string, 0, len(allowlist))
for m := range allowlist {
merged.ModelAllowlist = append(merged.ModelAllowlist, m)
}
sort.Strings(merged.ModelAllowlist)
}
return merged
}
// mergeGuardrail folds a single guardrail's prompt-capture settings into the
// running merge: enabled / redact-pii stick once any enabling guardrail turns
// them on.
func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) {
// mergeGuardrail folds a single guardrail's enabled checks into the
// running merge: model-allowlist models join the shared set (and flip
// allowlistEnabled), and prompt-capture / redact-pii stick once any
// enabling guardrail turns them on.
//
// TokenLimits, Budget, and Retention have moved off guardrails — token
// and budget caps now live on the Policy itself (Policy.Limits) and
// retention moves to account-level Settings — so they are not merged here.
func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails, allowlist map[string]struct{}, allowlistEnabled *bool) {
if g.Checks.ModelAllowlist.Enabled {
*allowlistEnabled = true
for _, m := range g.Checks.ModelAllowlist.Models {
if m != "" {
allowlist[m] = struct{}{}
}
}
}
if g.Checks.PromptCapture.Enabled {
merged.PromptCapture.Enabled = true
if g.Checks.PromptCapture.RedactPii {

View File

@@ -81,8 +81,8 @@ func TestSynthesizeServices_RealStore_PromptCaptureAccountIsSoleControl(t *testi
require.Len(t, services, 1)
cfg := decodeServiceGuardrailConfig(t, services[0])
assert.Equal(t, map[string][]string{"prov-1": {"gpt-5.4"}}, cfg.ProviderAllowlists,
"model allowlist is a pure policy guardrail and must reach the per-provider config")
assert.Equal(t, []string{"gpt-5.4"}, cfg.ModelAllowlist,
"model allowlist is a pure policy guardrail and must always reach the config")
assert.False(t, cfg.PromptCapture.Enabled,
"prompt capture must be off when the account toggle is off, even with a capture-enabled guardrail")
}
@@ -172,7 +172,7 @@ func TestSynthesizeServices_RealStore_NoGuardrail_CaptureOff(t *testing.T) {
require.Len(t, services, 1, "exactly one synth service expected")
cfg := decodeServiceGuardrailConfig(t, services[0])
assert.Empty(t, cfg.ProviderAllowlists, "no guardrail → provider unrestricted (absent from map)")
assert.Empty(t, cfg.ModelAllowlist, "no guardrail → no allowlist")
assert.False(t, cfg.PromptCapture.Enabled, "no guardrail → prompt capture off by default")
assert.False(t, cfg.PromptCapture.RedactPii, "no guardrail → redact off by default")
}

View File

@@ -1,95 +0,0 @@
package agentnetwork
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
)
// policyForProviders builds an enabled policy authorising the given providers
// under the given guardrails (both optional). Groups are irrelevant to
// buildProviderAllowlists, which keys purely on destination provider.
func policyForProviders(id string, guardrailIDs []string, providerIDs ...string) *types.Policy {
return &types.Policy{
ID: id,
Enabled: true,
DestinationProviderIDs: providerIDs,
GuardrailIDs: guardrailIDs,
}
}
func TestBuildProviderAllowlists(t *testing.T) {
byID := map[string]*types.Guardrail{
"g-4o": allowlistGuardrail("g-4o", "acc-1", "gpt-4o"),
"g-opus": allowlistGuardrail("g-opus", "acc-1", "claude-opus-4"),
"g-disabled": {ID: "g-disabled", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}}},
}
t.Run("all authorising policies restrict yields per-provider union", func(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
policyForProviders("p2", []string{"g-opus"}, "prov-x"),
}
got := buildProviderAllowlists(policies, byID)
assert.Equal(t, map[string][]string{"prov-x": {"claude-opus-4", "gpt-4o"}}, got,
"a provider every policy restricts carries the sorted union of their models")
})
t.Run("any un-guardrailed policy leaves the provider unrestricted (omitted)", func(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
policyForProviders("p2", nil, "prov-x"), // no guardrail
}
got := buildProviderAllowlists(policies, byID)
assert.NotContains(t, got, "prov-x",
"a provider reachable by an un-guardrailed policy must be omitted so the proxy treats it as unrestricted")
})
t.Run("a disabled allowlist counts as unrestricted", func(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-disabled"}, "prov-x"),
}
got := buildProviderAllowlists(policies, byID)
assert.NotContains(t, got, "prov-x",
"a policy whose only guardrail has a disabled allowlist is unrestricted")
})
t.Run("providers are isolated from one another", func(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
policyForProviders("p2", []string{"g-opus"}, "prov-y"),
}
got := buildProviderAllowlists(policies, byID)
assert.Equal(t, []string{"gpt-4o"}, got["prov-x"], "prov-x keeps only its own model")
assert.Equal(t, []string{"claude-opus-4"}, got["prov-y"], "prov-y keeps only its own model")
})
t.Run("one policy authorising two providers restricts both", func(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o"}, "prov-x", "prov-y"),
}
got := buildProviderAllowlists(policies, byID)
assert.Equal(t, []string{"gpt-4o"}, got["prov-x"])
assert.Equal(t, []string{"gpt-4o"}, got["prov-y"])
})
t.Run("union across a single policy's guardrails", func(t *testing.T) {
policies := []*types.Policy{
policyForProviders("p1", []string{"g-4o", "g-opus"}, "prov-x"),
}
got := buildProviderAllowlists(policies, byID)
assert.ElementsMatch(t, []string{"claude-opus-4", "gpt-4o"}, got["prov-x"],
"a policy's own multiple allowlist guardrails union together")
})
t.Run("an enabled allowlist with no models denies everything", func(t *testing.T) {
empty := map[string]*types.Guardrail{"g-empty": allowlistGuardrail("g-empty", "acc-1")}
got := buildProviderAllowlists([]*types.Policy{
policyForProviders("p1", []string{"g-empty"}, "prov-x"),
}, empty)
assert.Equal(t, map[string][]string{"prov-x": {}}, got,
"an enabled-but-empty allowlist is restricted with an empty set, not unrestricted")
})
}

View File

@@ -1031,12 +1031,8 @@ func TestSynthesizeServices_GuardrailMerge_AllowlistUnion_LimitsRestrictive(t *t
var cfg guardrailConfig
require.NoError(t, json.Unmarshal(guardrailJSON, &cfg), "guardrail config must unmarshal cleanly")
// Both policies restrict the same provider, so the per-provider backstop
// carries the union of their models — a coarse gate that management's
// per-policy/group check narrows; it only blocks models outside the union
// when management is down.
assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ProviderAllowlists["prov-1"],
"per-provider allowlist union must keep both models")
assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ModelAllowlist,
"model allowlist union must keep both models")
}
func TestSynthesizeServices_BackfillsMissingSessionKeys(t *testing.T) {

View File

@@ -41,24 +41,8 @@ type AgentNetworkAccessLog struct {
InputTokens int64
OutputTokens int64
TotalTokens int64
// Prompt-cache buckets: read + write token counts.
CachedInputTokens int64
CacheCreationTokens int64
// Per-bucket cost breakdown — one column per token bucket the provider
// bills separately. These four are the only cost state stored: the total
// and the cache portion are derived on read (TotalCostUSD / CacheCostUSD)
// rather than stored alongside, so a stored aggregate can never drift out
// of step with the components it summarises.
//
// default:0 matters on upgrade: these columns are ALTER TABLE ADD COLUMN
// on an existing table, and without it every historical row holds NULL —
// which a raw SUM()/scan into float64 can't read. The default backfills
// them as 0, so pre-upgrade rows report an unknown split, not an error.
InputCostUSD float64 `gorm:"not null;default:0"`
CachedInputCostUSD float64 `gorm:"not null;default:0"`
CacheCreationCostUSD float64 `gorm:"not null;default:0"`
OutputCostUSD float64 `gorm:"not null;default:0"`
Stream bool
CostUSD float64
Stream bool
// Prompt capture. Only populated when prompt collection is enabled
// (account master switch AND policy guardrail). Heavy free text.
@@ -76,44 +60,19 @@ type AgentNetworkAccessLog struct {
// the reverse-proxy AccessLogEntry table.
func (AgentNetworkAccessLog) TableName() string { return "agent_network_access_log" }
// CostUSDSQLExpr is the SQL sum of the per-bucket cost columns — the total cost
// of a row. Used wherever a query has to sort or aggregate on total cost now
// that no cost_usd column is stored. Plain arithmetic over NOT NULL columns, so
// it stays portable across SQLite and Postgres.
const CostUSDSQLExpr = "(input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd)"
// TotalCostUSD is the request's total cost: the sum of the four per-bucket
// costs. Derived rather than stored so it cannot disagree with the breakdown.
func (a *AgentNetworkAccessLog) TotalCostUSD() float64 {
return a.InputCostUSD + a.CachedInputCostUSD + a.CacheCreationCostUSD + a.OutputCostUSD
}
// CacheCostUSD is the portion of the total billed for prompt-cache buckets:
// cache reads plus cache writes.
func (a *AgentNetworkAccessLog) CacheCostUSD() float64 {
return a.CachedInputCostUSD + a.CacheCreationCostUSD
}
// ToAPIResponse renders the flattened entry as the API representation.
func (a *AgentNetworkAccessLog) ToAPIResponse() api.AgentNetworkAccessLog {
out := api.AgentNetworkAccessLog{
Id: a.ID,
ServiceId: a.ServiceID,
Timestamp: a.Timestamp,
StatusCode: a.StatusCode,
DurationMs: int(a.Duration.Milliseconds()),
InputTokens: a.InputTokens,
OutputTokens: a.OutputTokens,
TotalTokens: a.TotalTokens,
CachedInputTokens: a.CachedInputTokens,
CacheCreationTokens: a.CacheCreationTokens,
InputCostUsd: a.InputCostUSD,
CachedInputCostUsd: a.CachedInputCostUSD,
CacheCreationCostUsd: a.CacheCreationCostUSD,
OutputCostUsd: a.OutputCostUSD,
CostUsd: a.TotalCostUSD(),
CacheCostUsd: a.CacheCostUSD(),
Stream: &a.Stream,
Id: a.ID,
ServiceId: a.ServiceID,
Timestamp: a.Timestamp,
StatusCode: a.StatusCode,
DurationMs: int(a.Duration.Milliseconds()),
InputTokens: a.InputTokens,
OutputTokens: a.OutputTokens,
TotalTokens: a.TotalTokens,
CostUsd: a.CostUSD,
Stream: &a.Stream,
}
out.UserId = strPtr(a.UserID)
@@ -153,36 +112,20 @@ func strPtr(s string) *string {
// summary plus its ordered entries. Assembled in Go from a page of entries — it
// is not a stored table.
type AgentNetworkAccessLogSession struct {
SessionID string // empty for a session-less (singleton) request
UserID string
GroupIDs []string // union of the entries' authorising groups
StartedAt time.Time
EndedAt time.Time
RequestCount int
InputTokens int64
OutputTokens int64
TotalTokens int64
CachedInputTokens int64
CacheCreationTokens int64
InputCostUSD float64
CachedInputCostUSD float64
CacheCreationCostUSD float64
OutputCostUSD float64
Providers []string // distinct vendors seen in the session
Models []string // distinct models seen in the session
Decision string // "deny" if any entry was denied, else "allow"
Entries []*AgentNetworkAccessLog
}
// TotalCostUSD is the session's total cost: the sum of the four per-bucket
// costs accumulated across its entries.
func (sess *AgentNetworkAccessLogSession) TotalCostUSD() float64 {
return sess.InputCostUSD + sess.CachedInputCostUSD + sess.CacheCreationCostUSD + sess.OutputCostUSD
}
// CacheCostUSD is the session's prompt-cache spend: cache reads plus writes.
func (sess *AgentNetworkAccessLogSession) CacheCostUSD() float64 {
return sess.CachedInputCostUSD + sess.CacheCreationCostUSD
SessionID string // empty for a session-less (singleton) request
UserID string
GroupIDs []string // union of the entries' authorising groups
StartedAt time.Time
EndedAt time.Time
RequestCount int
InputTokens int64
OutputTokens int64
TotalTokens int64
CostUSD float64
Providers []string // distinct vendors seen in the session
Models []string // distinct models seen in the session
Decision string // "deny" if any entry was denied, else "allow"
Entries []*AgentNetworkAccessLog
}
// sessionKey is the grouping key for an entry: its session id, or — when the
@@ -262,12 +205,7 @@ func (sess *AgentNetworkAccessLogSession) foldEntry(sk *sessionSeen, e *AgentNet
sess.InputTokens += e.InputTokens
sess.OutputTokens += e.OutputTokens
sess.TotalTokens += e.TotalTokens
sess.CachedInputTokens += e.CachedInputTokens
sess.CacheCreationTokens += e.CacheCreationTokens
sess.InputCostUSD += e.InputCostUSD
sess.CachedInputCostUSD += e.CachedInputCostUSD
sess.CacheCreationCostUSD += e.CacheCreationCostUSD
sess.OutputCostUSD += e.OutputCostUSD
sess.CostUSD += e.CostUSD
if e.Timestamp.Before(sess.StartedAt) {
sess.StartedAt = e.Timestamp
}
@@ -310,22 +248,15 @@ func (sess *AgentNetworkAccessLogSession) ToAPIResponse() api.AgentNetworkAccess
}
out := api.AgentNetworkAccessLogSession{
StartedAt: sess.StartedAt,
EndedAt: sess.EndedAt,
RequestCount: sess.RequestCount,
InputTokens: sess.InputTokens,
OutputTokens: sess.OutputTokens,
TotalTokens: sess.TotalTokens,
CachedInputTokens: sess.CachedInputTokens,
CacheCreationTokens: sess.CacheCreationTokens,
InputCostUsd: sess.InputCostUSD,
CachedInputCostUsd: sess.CachedInputCostUSD,
CacheCreationCostUsd: sess.CacheCreationCostUSD,
OutputCostUsd: sess.OutputCostUSD,
CostUsd: sess.TotalCostUSD(),
CacheCostUsd: sess.CacheCostUSD(),
Decision: sess.Decision,
Entries: entries,
StartedAt: sess.StartedAt,
EndedAt: sess.EndedAt,
RequestCount: sess.RequestCount,
InputTokens: sess.InputTokens,
OutputTokens: sess.OutputTokens,
TotalTokens: sess.TotalTokens,
CostUsd: sess.CostUSD,
Decision: sess.Decision,
Entries: entries,
}
out.SessionId = strPtr(sess.SessionID)
out.UserId = strPtr(sess.UserID)

View File

@@ -54,7 +54,7 @@ var accessLogSortFields = map[string]string{
"provider": "provider",
"status_code": "status_code",
"duration": "duration",
"cost_usd": CostUSDSQLExpr,
"cost_usd": "cost_usd",
"total_tokens": "total_tokens",
"user_id": "user_id",
"decision": "decision",
@@ -70,7 +70,7 @@ var accessLogSortFields = map[string]string{
var sessionSortExprs = map[string]string{ //nolint:gosec // G101 false positive: "total_tokens" sort key, not a credential
"timestamp": "MAX(timestamp)",
"started_at": "MIN(timestamp)",
"cost_usd": "SUM" + CostUSDSQLExpr,
"cost_usd": "SUM(cost_usd)",
"total_tokens": "SUM(total_tokens)",
"duration": "SUM(duration)",
"request_count": "COUNT(*)",

View File

@@ -1,124 +0,0 @@
package types
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// costRow builds an access-log entry carrying only a cost breakdown — the rest
// of the row is irrelevant to the summation identities under test.
func costRow(id, session string, ts time.Time, in, cachedIn, cacheCreate, out float64) *AgentNetworkAccessLog {
return &AgentNetworkAccessLog{
ID: id,
SessionID: session,
Timestamp: ts,
InputCostUSD: in,
CachedInputCostUSD: cachedIn,
CacheCreationCostUSD: cacheCreate,
OutputCostUSD: out,
}
}
// TestAPIResponse_CostComponentsSumToAggregates is the contract a client adding
// up an API response depends on: within a single rendered object, the four
// per-bucket fields sum to cost_usd, and the two cache fields sum to
// cache_cost_usd. Uses rates that are not exactly representable in binary
// floating point, so the identity is checked against real arithmetic rather
// than round numbers.
func TestAPIResponse_CostComponentsSumToAggregates(t *testing.T) {
row := costRow("r1", "s1", time.Now(), 0.000768, 0.0002304, 0.00192, 0.003)
api := row.ToAPIResponse()
assert.InDelta(t, api.InputCostUsd+api.CachedInputCostUsd+api.CacheCreationCostUsd+api.OutputCostUsd,
api.CostUsd, 1e-12, "rendered buckets must sum to the rendered cost_usd")
assert.InDelta(t, api.CachedInputCostUsd+api.CacheCreationCostUsd, api.CacheCostUsd, 1e-12,
"rendered cache buckets must sum to the rendered cache_cost_usd")
assert.InDelta(t, 0.0059184, api.CostUsd, 1e-12, "total is the exact sum, not a separately rounded figure")
assert.InDelta(t, 0.0021504, api.CacheCostUsd, 1e-12, "cache cost is the exact sum of the two cache buckets")
}
// TestSessionSummary_SumsMatchSummedEntries proves a session summary equals the
// sum of the entries it renders: a client that adds up the entries itself must
// land on the same number the summary reports, per bucket and in total.
func TestSessionSummary_SumsMatchSummedEntries(t *testing.T) {
base := time.Date(2026, 5, 5, 10, 0, 0, 0, time.UTC)
entries := []*AgentNetworkAccessLog{
costRow("r1", "s1", base, 0.000768, 0.0002304, 0.00192, 0.003),
costRow("r2", "s1", base.Add(time.Minute), 0.000625, 0.0009375, 0, 0.005),
costRow("r3", "s1", base.Add(2*time.Minute), 0.0000016, 0, 0, 0.0000032),
}
sessions := FoldAccessLogSessions([]string{"s1"}, entries)
require.Len(t, sessions, 1)
sess := sessions[0].ToAPIResponse()
var wantInput, wantCachedInput, wantCacheCreation, wantOutput float64
for _, e := range entries {
wantInput += e.InputCostUSD
wantCachedInput += e.CachedInputCostUSD
wantCacheCreation += e.CacheCreationCostUSD
wantOutput += e.OutputCostUSD
}
assert.InDelta(t, wantInput, sess.InputCostUsd, 1e-12, "session input cost is the sum of its entries")
assert.InDelta(t, wantCachedInput, sess.CachedInputCostUsd, 1e-12, "session cache-read cost is the sum of its entries")
assert.InDelta(t, wantCacheCreation, sess.CacheCreationCostUsd, 1e-12, "session cache-write cost is the sum of its entries")
assert.InDelta(t, wantOutput, sess.OutputCostUsd, 1e-12, "session output cost is the sum of its entries")
assert.InDelta(t, wantInput+wantCachedInput+wantCacheCreation+wantOutput, sess.CostUsd, 1e-12,
"session total equals the summed entry buckets")
// Summing the rendered entries must give the same answer as reading the
// summary — the property a UI relies on when it totals a table itself.
var fromEntries float64
for _, e := range sess.Entries {
fromEntries += e.CostUsd
}
assert.InDelta(t, sess.CostUsd, fromEntries, 1e-12, "summary total must match the summed rendered entries")
// The sub-microdollar row must still contribute; it would vanish under
// 6-decimal quantisation.
assert.Greater(t, sess.InputCostUsd, 0.001393, "small-cost rows must not be quantised away")
}
// TestUsageBuckets_SumsMatchSummedRows proves the same identity one level up:
// a usage bucket equals the sum of the ledger rows folded into it, and the
// buckets together equal the whole range.
func TestUsageBuckets_SumsMatchSummedRows(t *testing.T) {
day1 := time.Date(2026, 5, 5, 9, 0, 0, 0, time.UTC)
day2 := time.Date(2026, 5, 6, 9, 0, 0, 0, time.UTC)
rows := []*AgentNetworkUsage{
{ID: "u1", Timestamp: day1, InputCostUSD: 0.000768, CachedInputCostUSD: 0.0002304, CacheCreationCostUSD: 0.00192, OutputCostUSD: 0.003},
{ID: "u2", Timestamp: day1.Add(time.Hour), InputCostUSD: 0.000625, CachedInputCostUSD: 0.0009375, OutputCostUSD: 0.005},
{ID: "u3", Timestamp: day2, InputCostUSD: 0.0000016, OutputCostUSD: 0.0000032},
}
buckets := AggregateUsageByGranularity(rows, UsageGranularityDay)
require.Len(t, buckets, 2, "two distinct days expected")
var total, cache float64
for _, b := range buckets {
api := b.ToAPIResponse()
assert.InDelta(t, api.InputCostUsd+api.CachedInputCostUsd+api.CacheCreationCostUsd+api.OutputCostUsd,
api.CostUsd, 1e-12, "each bucket's components must sum to its cost_usd")
total += api.CostUsd
cache += api.CacheCostUsd
}
var wantTotal, wantCache float64
for _, r := range rows {
wantTotal += r.TotalCostUSD()
wantCache += r.CacheCostUSD()
}
assert.InDelta(t, wantTotal, total, 1e-12, "buckets must sum to the total across all ledger rows")
assert.InDelta(t, wantCache, cache, 1e-12, "buckets must sum to the cache spend across all ledger rows")
// A month bucket over the same rows must total identically — regrouping
// changes the partition, never the sum.
monthly := AggregateUsageByGranularity(rows, UsageGranularityMonth)
require.Len(t, monthly, 1)
assert.InDelta(t, wantTotal, monthly[0].ToAPIResponse().CostUsd, 1e-12,
"re-bucketing at a different granularity must preserve the total")
}

View File

@@ -25,19 +25,8 @@ type AgentNetworkUsage struct {
InputTokens int64
OutputTokens int64
TotalTokens int64
// Prompt-cache buckets: read + write token counts.
CachedInputTokens int64
CacheCreationTokens int64
// Per-bucket cost breakdown, mirroring AgentNetworkAccessLog — the only
// cost state stored; total and cache portion are derived on read. Kept on
// the usage ledger too so spend can be attributed per bucket even for
// accounts with log collection turned off. See AgentNetworkAccessLog for
// why the columns carry a zero default.
InputCostUSD float64 `gorm:"not null;default:0"`
CachedInputCostUSD float64 `gorm:"not null;default:0"`
CacheCreationCostUSD float64 `gorm:"not null;default:0"`
OutputCostUSD float64 `gorm:"not null;default:0"`
CreatedAt time.Time
CostUSD float64
CreatedAt time.Time
}
// TableName keeps usage records in their own stripped table. Named
@@ -45,17 +34,6 @@ type AgentNetworkUsage struct {
// agent_network_usage table in a shared database.
func (AgentNetworkUsage) TableName() string { return "agent_network_request_usage" }
// TotalCostUSD is the request's total cost: the sum of the four per-bucket
// costs. Derived rather than stored so it cannot disagree with the breakdown.
func (u *AgentNetworkUsage) TotalCostUSD() float64 {
return u.InputCostUSD + u.CachedInputCostUSD + u.CacheCreationCostUSD + u.OutputCostUSD
}
// CacheCostUSD is the portion of the total billed for prompt-cache buckets.
func (u *AgentNetworkUsage) CacheCostUSD() float64 {
return u.CachedInputCostUSD + u.CacheCreationCostUSD
}
// AgentNetworkUsageGroup is the normalised many-to-many row linking a usage
// record to one authorising group, mirroring AgentNetworkAccessLogGroup so the
// usage overview can filter by group with a `group_id IN (...)` join.

View File

@@ -33,45 +33,21 @@ func ParseUsageGranularity(s string) UsageGranularity {
// AgentNetworkUsageBucket is one aggregated usage time bucket. PeriodStart is
// the UTC start of the bucket as YYYY-MM-DD.
type AgentNetworkUsageBucket struct {
PeriodStart string
InputTokens int64
OutputTokens int64
TotalTokens int64
CachedInputTokens int64
CacheCreationTokens int64
InputCostUSD float64
CachedInputCostUSD float64
CacheCreationCostUSD float64
OutputCostUSD float64
}
// TotalCostUSD is the bucket's total spend: the sum of the four per-bucket
// costs. Derived rather than accumulated separately so it cannot disagree with
// the components.
func (b *AgentNetworkUsageBucket) TotalCostUSD() float64 {
return b.InputCostUSD + b.CachedInputCostUSD + b.CacheCreationCostUSD + b.OutputCostUSD
}
// CacheCostUSD is the bucket's prompt-cache spend: cache reads plus writes.
func (b *AgentNetworkUsageBucket) CacheCostUSD() float64 {
return b.CachedInputCostUSD + b.CacheCreationCostUSD
PeriodStart string
InputTokens int64
OutputTokens int64
TotalTokens int64
CostUSD float64
}
// ToAPIResponse renders the bucket as the API representation.
func (b *AgentNetworkUsageBucket) ToAPIResponse() api.AgentNetworkUsageBucket {
return api.AgentNetworkUsageBucket{
PeriodStart: b.PeriodStart,
InputTokens: b.InputTokens,
OutputTokens: b.OutputTokens,
TotalTokens: b.TotalTokens,
CachedInputTokens: b.CachedInputTokens,
CacheCreationTokens: b.CacheCreationTokens,
InputCostUsd: b.InputCostUSD,
CachedInputCostUsd: b.CachedInputCostUSD,
CacheCreationCostUsd: b.CacheCreationCostUSD,
OutputCostUsd: b.OutputCostUSD,
CostUsd: b.TotalCostUSD(),
CacheCostUsd: b.CacheCostUSD(),
PeriodStart: b.PeriodStart,
InputTokens: b.InputTokens,
OutputTokens: b.OutputTokens,
TotalTokens: b.TotalTokens,
CostUsd: b.CostUSD,
}
}
@@ -108,12 +84,7 @@ func AggregateUsageByGranularity(rows []*AgentNetworkUsage, g UsageGranularity)
b.InputTokens += r.InputTokens
b.OutputTokens += r.OutputTokens
b.TotalTokens += r.TotalTokens
b.CachedInputTokens += r.CachedInputTokens
b.CacheCreationTokens += r.CacheCreationTokens
b.InputCostUSD += r.InputCostUSD
b.CachedInputCostUSD += r.CachedInputCostUSD
b.CacheCreationCostUSD += r.CacheCreationCostUSD
b.OutputCostUSD += r.OutputCostUSD
b.CostUSD += r.CostUSD
}
out := make([]*AgentNetworkUsageBucket, 0, len(byPeriod))

View File

@@ -20,17 +20,15 @@ import (
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/keepalive"
cachestore "github.com/eko/gocache/lib/v4/store"
"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"
@@ -70,8 +68,8 @@ func (s *BaseServer) Metrics() telemetry.AppMetrics {
// CacheStore returns a shared cache store backed by Redis or in-memory depending on the environment.
// All consumers should reuse this store to avoid creating multiple Redis connections.
func (s *BaseServer) CacheStore() cachestore.StoreInterface {
return Create(s, func() cachestore.StoreInterface {
func (s *BaseServer) CacheStore() nbcache.Store {
return Create(s, func() nbcache.Store {
cs, err := nbcache.NewStore(context.Background(), nbcache.DefaultStoreMaxTimeout, nbcache.DefaultStoreCleanupInterval, nbcache.DefaultStoreMaxConn)
if err != nil {
log.Fatalf("failed to create shared cache store: %v", err)

View File

@@ -5,6 +5,9 @@ import (
"strconv"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/networkmap"
@@ -163,7 +166,7 @@ func (e *componentEncoder) indexAllPeers() {
}
}
func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
func (e *componentEncoder) appendPeer(p *nbpeer.Peer) uint32 {
if idx, ok := e.peerOrder[p.ID]; ok {
return idx
}
@@ -177,7 +180,7 @@ func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
// (c.RouterPeers may contain peers not in c.Peers when validation rules drop
// them) and returns their wire indexes for the RouterPeerIndexes field. Must
// run before any encoder that resolves peer ids via e.peerOrder.
func (e *componentEncoder) indexRouterPeers(routers map[string]*types.ComponentPeer) []uint32 {
func (e *componentEncoder) indexRouterPeers(routers map[string]*nbpeer.Peer) []uint32 {
if len(routers) == 0 {
return nil
}
@@ -511,7 +514,7 @@ func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
return out
}
func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentResource) []*proto.NetworkResourceRaw {
func (e *componentEncoder) encodeNetworkResources(resources []*resourceTypes.NetworkResource) []*proto.NetworkResourceRaw {
if len(resources) == 0 {
return nil
}
@@ -540,7 +543,7 @@ func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentRe
return out
}
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*types.ComponentRouter) map[string]*proto.NetworkRouterList {
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*routerTypes.NetworkRouter) map[string]*proto.NetworkRouterList {
if len(routersMap) == 0 {
return nil
}
@@ -689,20 +692,20 @@ func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
return out
}
func toPeerCompact(p *types.ComponentPeer) *proto.PeerCompact {
func toPeerCompact(p *nbpeer.Peer) *proto.PeerCompact {
pc := &proto.PeerCompact{
WgPubKey: decodeWgKey(p.Key),
SshPubKey: []byte(p.SSHKey),
DnsLabel: p.DNSLabel,
AgentVersion: p.AgentVersion,
AddedWithSsoLogin: p.AddedWithSSOLogin,
AgentVersion: p.Meta.WtVersion,
AddedWithSsoLogin: p.UserID != "",
LoginExpirationEnabled: p.LoginExpirationEnabled,
SshEnabled: p.SSHEnabled,
SupportsIpv6: p.SupportsIPv6,
SupportsSourcePrefixes: p.SupportsSourcePrefixes,
ServerSshAllowed: p.ServerSSHAllowed,
SupportsIpv6: p.SupportsIPv6(),
SupportsSourcePrefixes: p.SupportsSourcePrefixes(),
ServerSshAllowed: p.Meta.Flags.ServerSSHAllowed,
}
if !p.LastLogin.IsZero() {
if p.LastLogin != nil {
pc.LastLoginUnixNano = p.LastLogin.UnixNano()
}
switch {

View File

@@ -15,6 +15,9 @@ import (
goproto "google.golang.org/protobuf/proto"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/proto"
@@ -152,28 +155,29 @@ func envelopesEquivalent(a, b *proto.NetworkMapEnvelope) bool {
}
func newTestComponents() *types.NetworkMapComponents {
peerA := &types.ComponentPeer{
ID: "peer-a",
Key: testWgKeyA,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peera",
SSHKey: "ssh-a",
AgentVersion: "0.40.0",
peerA := &nbpeer.Peer{
ID: "peer-a",
Key: testWgKeyA,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peera",
SSHKey: "ssh-a",
Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now()},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
peerB := &types.ComponentPeer{
ID: "peer-b",
Key: testWgKeyB,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
DNSLabel: "peerb",
AgentVersion: "0.25.0",
peerB := &nbpeer.Peer{
ID: "peer-b",
Key: testWgKeyB,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
DNSLabel: "peerb",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.25.0"},
}
peerC := &types.ComponentPeer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
AgentVersion: "0.40.0",
peerC := &nbpeer.Peer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
return &types.NetworkMapComponents{
@@ -187,12 +191,12 @@ func newTestComponents() *types.NetworkMapComponents {
PeerLoginExpirationEnabled: true,
PeerLoginExpiration: 2 * time.Hour,
},
Peers: map[string]*types.ComponentPeer{
Peers: map[string]*nbpeer.Peer{
"peer-a": peerA,
"peer-b": peerB,
"peer-c": peerC,
},
Groups: map[string]*types.ComponentGroup{
Groups: map[string]*types.Group{
"group-src": {ID: "group-src", PublicID: "1", Name: "Src", Peers: []string{"peer-a"}},
"group-dst": {ID: "group-dst", PublicID: "2", Name: "Dst", Peers: []string{"peer-b", "peer-c"}},
},
@@ -211,7 +215,7 @@ func newTestComponents() *types.NetworkMapComponents {
}},
},
},
RouterPeers: map[string]*types.ComponentPeer{"peer-c": peerC},
RouterPeers: map[string]*nbpeer.Peer{"peer-c": peerC},
}
}
@@ -377,12 +381,12 @@ func TestEncodeNetworkMapEnvelope_MalformedWgKey(t *testing.T) {
func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
c := newTestComponents()
v6Only := &types.ComponentPeer{
ID: "peer-v6",
Key: testWgKeyA,
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
DNSLabel: "peerv6",
AgentVersion: "0.40.0",
v6Only := &nbpeer.Peer{
ID: "peer-v6",
Key: testWgKeyA,
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
DNSLabel: "peerv6",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.Peers["peer-v6"] = v6Only
@@ -401,11 +405,11 @@ func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
c := newTestComponents()
c.Peers["peer-noip"] = &types.ComponentPeer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
AgentVersion: "0.40.0",
c.Peers["peer-noip"] = &nbpeer.Peer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -440,9 +444,9 @@ func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
c := newTestComponents()
now := time.Date(2024, 1, 2, 3, 4, 5, 0, time.UTC)
c.Peers["peer-a"].AddedWithSSOLogin = true
c.Peers["peer-a"].UserID = "user-1"
c.Peers["peer-a"].LoginExpirationEnabled = true
c.Peers["peer-a"].LastLogin = now
c.Peers["peer-a"].LastLogin = &now
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -553,7 +557,7 @@ func TestEncodeNetworkMapEnvelope_ResourceOnlyPolicyShippedAndIndexed(t *testing
}
// Resource must appear in components.NetworkResources with a seq id —
// encoder uses that to translate the xid map key to uint32.
c.NetworkResources = []*types.ComponentResource{
c.NetworkResources = []*resourceTypes.NetworkResource{
{ID: "resource-x", PublicID: "77", Name: "res-x", Enabled: true},
}
@@ -621,11 +625,11 @@ func TestEncodeNetworkMapEnvelope_PostureFailedPeers(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RoutersMap(t *testing.T) {
c := newTestComponents()
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
c.RoutersMap = map[string]map[string]*routerTypes.NetworkRouter{
"net-1": {
"peer-c": {
PublicID: "200",
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
ID: "router-1", PublicID: "200",
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
},
},
}
@@ -651,14 +655,14 @@ func TestEncodeNetworkMapEnvelope_RouterPeerNotInComponentsPeers(t *testing.T) {
// peer_index reference must still resolve.
c := newTestComponents()
delete(c.Peers, "peer-c")
routerPeer := &types.ComponentPeer{
routerPeer := &nbpeer.Peer{
ID: "peer-c", Key: testWgKeyC, IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc", AgentVersion: "0.40.0",
DNSLabel: "peerc", Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.RouterPeers = map[string]*types.ComponentPeer{"peer-c": routerPeer}
c.RouterPeers = map[string]*nbpeer.Peer{"peer-c": routerPeer}
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
"net-1": {"peer-c": {PublicID: "1", Peer: "peer-c", Enabled: true}},
c.RoutersMap = map[string]map[string]*routerTypes.NetworkRouter{
"net-1": {"peer-c": {ID: "r-1", PublicID: "1", Peer: "peer-c", Enabled: true}},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -691,9 +695,9 @@ func TestToProxyPatch_EmptyInputReturnsNil(t *testing.T) {
func TestToProxyPatch_PopulatesAllFields(t *testing.T) {
nm := &types.NetworkMap{
Peers: []*types.ComponentPeer{{
Peers: []*nbpeer.Peer{{
ID: "ext-peer", Key: testWgKeyA, IP: netip.AddrFrom4([4]byte{100, 64, 0, 9}),
DNSLabel: "extpeer", AgentVersion: "0.40.0",
DNSLabel: "extpeer", Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}},
FirewallRules: []*types.FirewallRule{{
PeerIP: "100.64.0.9", Action: "accept", Direction: 0, Protocol: "tcp",
@@ -776,6 +780,6 @@ 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]*nbpeer.Peer{"peer-id": {}}},
)
}

View File

@@ -19,10 +19,10 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
mkComponents := func(rule *types.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nbpeer.Peer) {
peer := &nbpeer.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
group := &types.ComponentGroup{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
group := &types.Group{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
return &types.NetworkMapComponents{
Peers: map[string]*types.ComponentPeer{targetPeerID: peer.ToComponent()},
Groups: map[string]*types.ComponentGroup{targetGroupID: group},
Peers: map[string]*nbpeer.Peer{targetPeerID: peer},
Groups: map[string]*types.Group{targetGroupID: group},
Policies: []*types.Policy{{
ID: "p",
Enabled: true,
@@ -158,8 +158,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
peer := &nbpeer.Peer{ID: "missing", SSHEnabled: true}
c := &types.NetworkMapComponents{
Peers: map[string]*types.ComponentPeer{}, // target peer NOT present
Groups: map[string]*types.ComponentGroup{
Peers: map[string]*nbpeer.Peer{}, // target peer NOT present
Groups: map[string]*types.Group{
"g": {ID: "g", Peers: []string{"missing"}},
},
Policies: []*types.Policy{{

View File

@@ -285,7 +285,6 @@ func (s *ProxyServiceServer) CheckLLMPolicyLimits(ctx context.Context, req *prot
UserID: req.GetUserId(),
GroupIDs: req.GetGroupIds(),
ProviderID: req.GetProviderId(),
Model: req.GetModel(),
})
if err != nil {
log.WithContext(ctx).Errorf("select policy for request: %v", err)

View File

@@ -1,138 +0,0 @@
package grpc
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc/codes"
grpcstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/shared/management/proto"
)
// fakeAgentNetworkLimits records the PolicySelectionInput it was invoked with
// and returns a pre-programmed result, so tests can assert what the handler
// forwards to the selector.
type fakeAgentNetworkLimits struct {
gotInput agentnetwork.PolicySelectionInput
result *agentnetwork.PolicySelectionResult
err error
}
func (f *fakeAgentNetworkLimits) SelectPolicyForRequest(_ context.Context, in agentnetwork.PolicySelectionInput) (*agentnetwork.PolicySelectionResult, error) {
f.gotInput = in
if f.err != nil {
return nil, f.err
}
return f.result, nil
}
func (f *fakeAgentNetworkLimits) RecordUsage(_ context.Context, _ agentnetwork.RecordUsageInput) error {
return nil
}
// TestCheckLLMPolicyLimits_ForwardsModelToSelector proves the wiring added here:
// the model the proxy extracted must reach the selector's Model unchanged,
// alongside the account/user/group/provider fields.
func TestCheckLLMPolicyLimits_ForwardsModelToSelector(t *testing.T) {
fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{Allow: true, SelectedPolicyID: "pol-1"}}
s := &ProxyServiceServer{}
s.SetAgentNetworkLimitsService(fake)
req := &proto.CheckLLMPolicyLimitsRequest{
AccountId: "acc-1",
UserId: "user-1",
GroupIds: []string{"grp-a", "grp-b"},
ProviderId: "prov-1",
Model: "claude-opus-4",
}
resp, err := s.CheckLLMPolicyLimits(context.Background(), req)
require.NoError(t, err)
require.NotNil(t, resp)
assert.Equal(t, "acc-1", fake.gotInput.AccountID)
assert.Equal(t, "user-1", fake.gotInput.UserID)
assert.Equal(t, []string{"grp-a", "grp-b"}, fake.gotInput.GroupIDs)
assert.Equal(t, "prov-1", fake.gotInput.ProviderID)
assert.Equal(t, "claude-opus-4", fake.gotInput.Model,
"the request's model must be forwarded to the selector")
}
// TestCheckLLMPolicyLimits_EmptyModelForwardedAsEmpty proves an undetermined
// model (empty string) is forwarded as-is; the selector decides how to treat it.
func TestCheckLLMPolicyLimits_EmptyModelForwardedAsEmpty(t *testing.T) {
fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{Allow: true}}
s := &ProxyServiceServer{}
s.SetAgentNetworkLimitsService(fake)
req := &proto.CheckLLMPolicyLimitsRequest{
AccountId: "acc-1",
ProviderId: "prov-1",
}
_, err := s.CheckLLMPolicyLimits(context.Background(), req)
require.NoError(t, err)
assert.Equal(t, "", fake.gotInput.Model, "an absent model must be forwarded as empty, not fabricated")
}
// TestCheckLLMPolicyLimits_DenyResponseCarriesModelBlockedCode proves the deny
// envelope surfaces the model-allowlist deny code + reason through the response.
func TestCheckLLMPolicyLimits_DenyResponseCarriesModelBlockedCode(t *testing.T) {
fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{
Allow: false,
DenyCode: "llm_policy.model_blocked",
DenyReason: `model "claude-opus-4" is not permitted by any applicable policy allowlist`,
}}
s := &ProxyServiceServer{}
s.SetAgentNetworkLimitsService(fake)
resp, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{
AccountId: "acc-1",
ProviderId: "prov-1",
Model: "claude-opus-4",
})
require.NoError(t, err)
require.NotNil(t, resp)
assert.Equal(t, "deny", resp.Decision)
assert.Equal(t, "llm_policy.model_blocked", resp.DenyCode)
assert.NotEmpty(t, resp.DenyReason)
assert.Empty(t, resp.SelectedPolicyId, "a denied request must carry no selected policy")
}
// TestCheckLLMPolicyLimits_SelectorErrorSurfacesAsInternal proves a selector
// failure surfaces as an Internal gRPC error rather than a silent allow.
func TestCheckLLMPolicyLimits_SelectorErrorSurfacesAsInternal(t *testing.T) {
fake := &fakeAgentNetworkLimits{err: errors.New("boom")}
s := &ProxyServiceServer{}
s.SetAgentNetworkLimitsService(fake)
_, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{
AccountId: "acc-1",
ProviderId: "prov-1",
Model: "gpt-4o",
})
require.Error(t, err)
st, ok := grpcstatus.FromError(err)
require.True(t, ok)
assert.Equal(t, codes.Internal, st.Code(), "selector errors must never fail open on the hot path")
}
// TestCheckLLMPolicyLimits_UnconfiguredServiceReturnsUnimplemented locks the
// fallback: with no limits service wired the RPC returns Unimplemented.
func TestCheckLLMPolicyLimits_UnconfiguredServiceReturnsUnimplemented(t *testing.T) {
s := &ProxyServiceServer{}
_, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{
AccountId: "acc-1",
ProviderId: "prov-1",
})
require.Error(t, err)
st, ok := grpcstatus.FromError(err)
require.True(t, ok)
assert.Equal(t, codes.Unimplemented, st.Code())
}

View File

@@ -7,9 +7,6 @@ import (
"errors"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/cache"
"github.com/eko/gocache/lib/v4/store"
)
const (
@@ -22,12 +19,17 @@ var (
ErrTokenExpired = errors.New("JWT expired")
)
type SessionStore struct {
cache *cache.Cache[string]
// TokenCache atomically records used JWTs until their expiration.
type TokenCache interface {
SetNX(ctx context.Context, key, value string, ttl time.Duration) (bool, error)
}
func NewSessionStore(cacheStore store.StoreInterface) *SessionStore {
return &SessionStore{cache: cache.New[string](cacheStore)}
type SessionStore struct {
cache TokenCache
}
func NewSessionStore(cacheStore TokenCache) *SessionStore {
return &SessionStore{cache: cacheStore}
}
// RegisterToken records a JWT until its exp time and rejects reuse.
@@ -38,19 +40,13 @@ func (s *SessionStore) RegisterToken(ctx context.Context, token string, expiresA
}
key := usedTokenKeyPrefix + hashToken(token)
_, err := s.cache.Get(ctx, key)
if err == nil {
return ErrTokenAlreadyUsed
}
var notFound *store.NotFound
if !errors.As(err, &notFound) {
return fmt.Errorf("failed to lookup used token entry: %w", err)
}
if err := s.cache.Set(ctx, key, usedTokenMarker, store.WithExpiration(ttl)); err != nil {
created, err := s.cache.SetNX(ctx, key, usedTokenMarker, ttl)
if err != nil {
return fmt.Errorf("failed to store used token entry: %w", err)
}
if !created {
return ErrTokenAlreadyUsed
}
return nil
}

View File

@@ -2,6 +2,7 @@ package auth
import (
"context"
"errors"
"testing"
"time"
@@ -38,6 +39,39 @@ func TestSessionStore_RegisterSameTokenTwiceIsRejected(t *testing.T) {
assert.ErrorIs(t, err, ErrTokenAlreadyUsed)
}
func TestSessionStore_ConcurrentRegistrationAllowsOneCaller(t *testing.T) {
s := newTestSessionStore(t)
ctx := context.Background()
const attempts = 100
start := make(chan struct{})
results := make(chan error, attempts)
for range attempts {
go func() {
<-start
results <- s.RegisterToken(ctx, "token", time.Now().Add(time.Hour))
}()
}
close(start)
succeeded := 0
alreadyUsed := 0
for range attempts {
err := <-results
switch {
case err == nil:
succeeded++
case errors.Is(err, ErrTokenAlreadyUsed):
alreadyUsed++
default:
require.NoError(t, err)
}
}
assert.Equal(t, 1, succeeded)
assert.Equal(t, attempts-1, alreadyUsed)
}
func TestSessionStore_RegisterDifferentTokensAreIndependent(t *testing.T) {
s := newTestSessionStore(t)
ctx := context.Background()
@@ -72,6 +106,23 @@ func TestSessionStore_EntryEvictsAtTTLAndAllowsReRegistration(t *testing.T) {
require.NoError(t, s.RegisterToken(ctx, token, time.Now().Add(time.Hour)))
}
type failingTokenCache struct {
err error
}
func (f failingTokenCache) SetNX(context.Context, string, string, time.Duration) (bool, error) {
return false, f.err
}
func TestSessionStore_CacheErrorIsReturned(t *testing.T) {
cacheErr := errors.New("cache unavailable")
s := NewSessionStore(failingTokenCache{err: cacheErr})
err := s.RegisterToken(context.Background(), "token", time.Now().Add(time.Hour))
require.Error(t, err)
assert.ErrorIs(t, err, cacheErr)
}
func TestHashToken_StableAndDoesNotLeak(t *testing.T) {
a := hashToken("tokenA")
b := hashToken("tokenB")

31
management/server/cache/memory.go vendored Normal file
View File

@@ -0,0 +1,31 @@
package cache
import (
"context"
"time"
"github.com/eko/gocache/lib/v4/store"
gocachestore "github.com/eko/gocache/store/go_cache/v4"
gocache "github.com/patrickmn/go-cache"
)
type goCacheStore struct {
store.StoreInterface
client *gocache.Cache
}
func newMemoryStore(maxTimeout, cleanupInterval time.Duration) Store {
client := gocache.New(maxTimeout, cleanupInterval)
return &goCacheStore{
StoreInterface: gocachestore.NewGoCache(client),
client: client,
}
}
func (s *goCacheStore) SetNX(_ context.Context, key, value string, ttl time.Duration) (bool, error) {
// Add only returns an error when a non-expired entry already exists.
if err := s.client.Add(key, value, ttl); err != nil {
return false, nil
}
return true, nil
}

49
management/server/cache/memory_test.go vendored Normal file
View File

@@ -0,0 +1,49 @@
package cache_test
import (
"context"
"testing"
"time"
"github.com/netbirdio/netbird/management/server/cache"
)
func TestMemoryStore(t *testing.T) {
memStore, err := cache.NewStore(context.Background(), 100*time.Millisecond, 300*time.Millisecond, 100)
if err != nil {
t.Fatalf("couldn't create memory store: %s", err)
}
ctx := context.Background()
key, value := "testing", "tested"
err = memStore.Set(ctx, key, value)
if err != nil {
t.Errorf("couldn't set testing data: %s", err)
}
result, err := memStore.Get(ctx, key)
if err != nil {
t.Errorf("couldn't get testing data: %s", err)
}
if value != result.(string) {
t.Errorf("value returned doesn't match testing data, got %s, expected %s", result, value)
}
created, err := memStore.SetNX(ctx, "atomic", value, 100*time.Millisecond)
if err != nil {
t.Fatalf("couldn't atomically set testing data: %s", err)
}
if !created {
t.Fatal("first atomic set should create the entry")
}
created, err = memStore.SetNX(ctx, "atomic", value, 100*time.Millisecond)
if err != nil {
t.Fatalf("couldn't atomically check testing data: %s", err)
}
if created {
t.Fatal("second atomic set should not replace the entry")
}
// test expiration
time.Sleep(300 * time.Millisecond)
_, err = memStore.Get(ctx, key)
if err == nil {
t.Error("value should not be found")
}
}

51
management/server/cache/redis.go vendored Normal file
View File

@@ -0,0 +1,51 @@
package cache
import (
"context"
"fmt"
"math"
"time"
"github.com/eko/gocache/lib/v4/store"
redisstore "github.com/eko/gocache/store/redis/v4"
"github.com/redis/go-redis/v9"
log "github.com/sirupsen/logrus"
)
type redisStore struct {
store.StoreInterface
client *redis.Client
}
func getRedisStore(ctx context.Context, redisEnvAddr string, maxConn int) (Store, error) {
options, err := redis.ParseURL(redisEnvAddr)
if err != nil {
return nil, fmt.Errorf("parsing redis cache url: %s", err)
}
options.MaxIdleConns = int(math.Ceil(float64(maxConn) * 0.5)) // 50% of max conns
options.MinIdleConns = int(math.Ceil(float64(maxConn) * 0.1)) // 10% of max conns
options.MaxActiveConns = maxConn
options.ConnMaxIdleTime = 30 * time.Minute
options.ConnMaxLifetime = 0
options.PoolTimeout = 10 * time.Second
redisClient := redis.NewClient(options)
subCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
defer cancel()
_, err = redisClient.Ping(subCtx).Result()
if err != nil {
return nil, err
}
log.WithContext(subCtx).Infof("using redis cache at %s", redisEnvAddr)
return &redisStore{
StoreInterface: redisstore.NewRedis(redisClient),
client: redisClient,
}, nil
}
func (s *redisStore) SetNX(ctx context.Context, key, value string, ttl time.Duration) (bool, error) {
return s.client.SetNX(ctx, key, value, ttl).Result()
}

View File

@@ -12,32 +12,6 @@ import (
"github.com/netbirdio/netbird/management/server/cache"
)
func TestMemoryStore(t *testing.T) {
memStore, err := cache.NewStore(context.Background(), 100*time.Millisecond, 300*time.Millisecond, 100)
if err != nil {
t.Fatalf("couldn't create memory store: %s", err)
}
ctx := context.Background()
key, value := "testing", "tested"
err = memStore.Set(ctx, key, value)
if err != nil {
t.Errorf("couldn't set testing data: %s", err)
}
result, err := memStore.Get(ctx, key)
if err != nil {
t.Errorf("couldn't get testing data: %s", err)
}
if value != result.(string) {
t.Errorf("value returned doesn't match testing data, got %s, expected %s", result, value)
}
// test expiration
time.Sleep(300 * time.Millisecond)
_, err = memStore.Get(ctx, key)
if err == nil {
t.Error("value should not be found")
}
}
func TestRedisStoreConnectionFailure(t *testing.T) {
t.Setenv(cache.RedisStoreEnvVar, "redis://127.0.0.1:6379")
_, err := cache.NewStore(context.Background(), 10*time.Millisecond, 30*time.Millisecond, 100)
@@ -94,6 +68,47 @@ func TestRedisStoreConnectionSuccess(t *testing.T) {
if value != r {
t.Errorf("value returned from redis doesn't match testing data, got %s, expected %s", r, value)
}
secondRedisStore, err := cache.NewStore(context.Background(), 100*time.Millisecond, 300*time.Millisecond, 100)
if err != nil {
t.Fatalf("couldn't create second redis store: %s", err)
}
start := make(chan struct{})
type setResult struct {
created bool
err error
}
results := make(chan setResult, 2)
for _, cacheStore := range []cache.Store{redisStore, secondRedisStore} {
go func() {
<-start
created, err := cacheStore.SetNX(ctx, "atomic", value, time.Second)
results <- setResult{created: created, err: err}
}()
}
close(start)
created := 0
for range 2 {
result := <-results
if result.err != nil {
t.Fatalf("atomic redis set failed: %s", result.err)
}
if result.created {
created++
}
}
if created != 1 {
t.Fatalf("expected exactly one redis client to create the entry, got %d", created)
}
ttl, err := redisClient.PTTL(ctx, "atomic").Result()
if err != nil {
t.Fatalf("couldn't read atomic entry TTL: %s", err)
}
if ttl <= 0 {
t.Fatalf("atomic entry should have a positive TTL, got %s", ttl)
}
// test expiration
time.Sleep(300 * time.Millisecond)
_, err = redisStore.Get(ctx, key)

View File

@@ -2,17 +2,10 @@ package cache
import (
"context"
"fmt"
"math"
"os"
"time"
"github.com/eko/gocache/lib/v4/store"
gocache_store "github.com/eko/gocache/store/go_cache/v4"
redis_store "github.com/eko/gocache/store/redis/v4"
gocache "github.com/patrickmn/go-cache"
"github.com/redis/go-redis/v9"
log "github.com/sirupsen/logrus"
)
// RedisStoreEnvVar is the environment variable that determines if a redis store should be used.
@@ -31,15 +24,21 @@ const (
DefaultStoreMaxConn = 1000
)
// Store extends the shared cache interface with atomic insertion support.
type Store interface {
store.StoreInterface
// SetNX atomically stores a value with a TTL only when the key does not exist.
SetNX(ctx context.Context, key, value string, ttl time.Duration) (bool, error)
}
// NewStore creates a new cache store with the given max timeout and cleanup interval. It checks for the environment Variable RedisStoreEnvVar
// to determine if a redis store should be used. If the environment variable is set, it will attempt to connect to the redis store.
func NewStore(ctx context.Context, maxTimeout, cleanupInterval time.Duration, maxConn int) (store.StoreInterface, error) {
func NewStore(ctx context.Context, maxTimeout, cleanupInterval time.Duration, maxConn int) (Store, error) {
redisAddr := GetAddrFromEnv()
if redisAddr != "" {
return getRedisStore(ctx, redisAddr, maxConn)
}
goc := gocache.New(maxTimeout, cleanupInterval)
return gocache_store.NewGoCache(goc), nil
return newMemoryStore(maxTimeout, cleanupInterval), nil
}
// GetAddrFromEnv returns the redis address from the environment variable RedisStoreEnvVar or its legacy counterpart.
@@ -50,29 +49,3 @@ func GetAddrFromEnv() string {
}
return addr
}
func getRedisStore(ctx context.Context, redisEnvAddr string, maxConn int) (store.StoreInterface, error) {
options, err := redis.ParseURL(redisEnvAddr)
if err != nil {
return nil, fmt.Errorf("parsing redis cache url: %s", err)
}
options.MaxIdleConns = int(math.Ceil(float64(maxConn) * 0.5)) // 50% of max conns
options.MinIdleConns = int(math.Ceil(float64(maxConn) * 0.1)) // 10% of max conns
options.MaxActiveConns = maxConn
options.ConnMaxIdleTime = 30 * time.Minute
options.ConnMaxLifetime = 0
options.PoolTimeout = 10 * time.Second
redisClient := redis.NewClient(options)
subCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
defer cancel()
_, err = redisClient.Ping(subCtx).Result()
if err != nil {
return nil, err
}
log.WithContext(subCtx).Infof("using redis cache at %s", redisEnvAddr)
return redis_store.NewRedis(redisClient), nil
}

View File

@@ -6,7 +6,6 @@ import (
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
@@ -31,10 +30,6 @@ type managerImpl struct {
accountManager account.Manager
}
func eventMetaResource(group *types.Group, resource *resourceTypes.NetworkResource) map[string]any {
return map[string]any{"name": group.Name, "id": group.ID, "resource_name": resource.Name, "resource_id": resource.ID, "resource_type": resource.Type}
}
type mockManager struct {
}
@@ -114,7 +109,7 @@ func (m *managerImpl) AddResourceToGroupInTransaction(ctx context.Context, trans
}
event := func() {
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, eventMetaResource(group, networkResource))
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, group.EventMetaResource(networkResource))
}
return event, nil
@@ -138,7 +133,7 @@ func (m *managerImpl) RemoveResourceFromGroupInTransaction(ctx context.Context,
}
event := func() {
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, eventMetaResource(group, networkResource))
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, group.EventMetaResource(networkResource))
}
return event, nil

View File

@@ -446,7 +446,7 @@ func (h *Handler) GetAccessiblePeers(w http.ResponseWriter, r *http.Request) {
netMap := account.GetPeerNetworkMapFromComponents(ctx, peerID, dns.CustomZone{}, nil, validPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(), nil, account.GetActiveGroupUsers())
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, account.Peers, dnsDomain))
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, dnsDomain))
}
func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request) {
@@ -534,20 +534,15 @@ func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request)
util.WriteJSONObject(r.Context(), w, resp)
}
// toAccessiblePeers rehydrates the calculated map's component peers into the
// account's full peer objects, which carry the location/status/meta fields
// the API response needs.
func toAccessiblePeers(netMap *types.NetworkMap, accountPeers map[string]*nbpeer.Peer, dnsDomain string) []api.AccessiblePeer {
func toAccessiblePeers(netMap *types.NetworkMap, dnsDomain string) []api.AccessiblePeer {
accessiblePeers := make([]api.AccessiblePeer, 0, len(netMap.Peers)+len(netMap.OfflinePeers))
add := func(peers []*types.ComponentPeer) {
for _, p := range peers {
if peer := accountPeers[p.ID]; peer != nil {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(peer, dnsDomain))
}
}
for _, p := range netMap.Peers {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
for _, p := range netMap.OfflinePeers {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
add(netMap.Peers)
add(netMap.OfflinePeers)
return accessiblePeers
}

View File

@@ -683,81 +683,3 @@ func BackfillPublicIDs[T any](ctx context.Context, db *gorm.DB) error {
log.WithContext(ctx).Infof("Backfill of empty public_id in table %s completed", tableName)
return nil
}
// FoldCostAggregatesIntoBuckets migrates a per-request cost table from the old
// "stored aggregate" shape (cost_usd + cache_cost_usd columns) to the per-bucket
// breakdown, where the total and cache portion are derived on read instead.
//
// The fold preserves both aggregates exactly for historical rows: the cache
// total moves into cached_input_cost_usd and the remainder into
// input_cost_usd, so a row's derived total and cache cost still match what it
// reported before the upgrade. The finer split is genuinely unknown for those
// rows — the old schema never recorded a read/write or input/output division —
// so it is lumped rather than guessed; only rows written after the upgrade
// carry a true four-way split.
//
// Dropping the columns before folding would zero every historical row's cost,
// so the update runs first and the drop only happens once it succeeds. A table
// with no cost_usd column has already been migrated (or was created fresh) and
// is skipped.
func FoldCostAggregatesIntoBuckets[T any](ctx context.Context, db *gorm.DB) error {
var model T
if !db.Migrator().HasTable(&model) {
log.WithContext(ctx).Debugf("table for %T does not exist, no cost-bucket migration needed", model)
return nil
}
if !db.Migrator().HasColumn(&model, "cost_usd") {
log.WithContext(ctx).Debugf("table for %T has no cost_usd column, cost buckets already migrated", model)
return nil
}
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(&model); err != nil {
return fmt.Errorf("parse model schema: %w", err)
}
tableName := stmt.Schema.Table
// COALESCE guards rows whose new columns were added as NULL by an earlier
// AutoMigrate run that predates the NOT NULL default.
hasCacheColumn := db.Migrator().HasColumn(&model, "cache_cost_usd")
cacheExpr := "0"
if hasCacheColumn {
cacheExpr = "COALESCE(cache_cost_usd, 0)"
}
if err := db.Transaction(func(tx *gorm.DB) error {
// Only touch rows that carry a legacy total and no breakdown yet, so
// the migration is idempotent and never overwrites a true split.
update := fmt.Sprintf(`UPDATE %s
SET input_cost_usd = COALESCE(cost_usd, 0) - %s,
cached_input_cost_usd = %s,
cache_creation_cost_usd = 0,
output_cost_usd = 0
WHERE COALESCE(cost_usd, 0) <> 0
AND COALESCE(input_cost_usd, 0) = 0
AND COALESCE(cached_input_cost_usd, 0) = 0
AND COALESCE(cache_creation_cost_usd, 0) = 0
AND COALESCE(output_cost_usd, 0) = 0`, tableName, cacheExpr, cacheExpr)
res := tx.Exec(update)
if res.Error != nil {
return fmt.Errorf("fold legacy cost aggregates in %s: %w", tableName, res.Error)
}
log.WithContext(ctx).Infof("folded legacy cost aggregates into per-bucket columns for %d rows in table %s", res.RowsAffected, tableName)
if err := tx.Migrator().DropColumn(&model, "cost_usd"); err != nil {
return fmt.Errorf("drop cost_usd from %s: %w", tableName, err)
}
if hasCacheColumn {
if err := tx.Migrator().DropColumn(&model, "cache_cost_usd"); err != nil {
return fmt.Errorf("drop cache_cost_usd from %s: %w", tableName, err)
}
}
return nil
}); err != nil {
return err
}
log.WithContext(ctx).Infof("migration of stored cost aggregates to per-bucket columns in table %s completed", tableName)
return nil
}

View File

@@ -16,7 +16,6 @@ import (
"gorm.io/driver/sqlite"
"gorm.io/gorm"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/migration"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/testutil"
@@ -640,99 +639,3 @@ func TestCleanupOrphanedResources_SkipsWhenForeignKeyExists(t *testing.T) {
db.Model(&testChildWithFK{}).Count(&count)
assert.Equal(t, int64(2), count, "Both rows should survive — migration must skip when FK constraint exists")
}
// legacyCostRow is the pre-breakdown shape of the usage table: cost was stored
// as a total plus a cache portion, with no per-bucket columns. Used to build a
// realistic pre-upgrade table for the fold migration to run against.
type legacyCostRow struct {
ID string `gorm:"primaryKey"`
AccountID string
Model string
CostUSD float64
CacheCostUSD float64
}
func (legacyCostRow) TableName() string { return "agent_network_request_usage" }
// TestFoldCostAggregatesIntoBuckets_PreservesHistoricalCost covers the upgrade
// path: a table written under the old schema must come out with its per-row
// total and cache cost unchanged, because dropping cost_usd without folding it
// forward would silently zero every historical row's spend.
func TestFoldCostAggregatesIntoBuckets_PreservesHistoricalCost(t *testing.T) {
ctx := context.Background()
db := setupDatabase(t)
// setupDatabase hands back a process-shared database, so start from a clean
// table rather than inheriting rows from another test.
require.NoError(t, db.Migrator().DropTable(&agentNetworkTypes.AgentNetworkUsage{}))
require.NoError(t, db.AutoMigrate(&legacyCostRow{}), "legacy table must be created")
require.NoError(t, db.Create(&legacyCostRow{
ID: "u1", AccountID: "acct-1", Model: "claude-sonnet-4-6", CostUSD: 0.0123, CacheCostUSD: 0.0029,
}).Error)
require.NoError(t, db.Create(&legacyCostRow{
ID: "u2", AccountID: "acct-1", Model: "gpt-4o", CostUSD: 0.5, CacheCostUSD: 0,
}).Error)
// A zero-cost row (denied / unpriced request) must stay zero, not be touched.
require.NoError(t, db.Create(&legacyCostRow{ID: "u3", AccountID: "acct-1", Model: "gw/unpriced"}).Error)
// AutoMigrate adds the per-bucket columns alongside the legacy ones, exactly
// as a real upgrade does before the post-auto migrations run.
require.NoError(t, db.AutoMigrate(&agentNetworkTypes.AgentNetworkUsage{}), "new columns must be added")
require.NoError(t, migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkUsage](ctx, db))
assert.False(t, db.Migrator().HasColumn(&agentNetworkTypes.AgentNetworkUsage{}, "cost_usd"),
"legacy cost_usd column must be dropped once folded")
assert.False(t, db.Migrator().HasColumn(&agentNetworkTypes.AgentNetworkUsage{}, "cache_cost_usd"),
"legacy cache_cost_usd column must be dropped once folded")
var rows []*agentNetworkTypes.AgentNetworkUsage
require.NoError(t, db.Order("id").Find(&rows).Error)
require.Len(t, rows, 3)
// u1: total and cache portion both preserved; the read/write and
// input/output splits are unknowable for a legacy row, so the cache total
// lands on cached_input and the remainder on input.
assert.InDelta(t, 0.0123, rows[0].TotalCostUSD(), 1e-9, "historical total must survive the fold")
assert.InDelta(t, 0.0029, rows[0].CacheCostUSD(), 1e-9, "historical cache cost must survive the fold")
assert.InDelta(t, 0.0094, rows[0].InputCostUSD, 1e-9, "non-cache remainder lands on input")
assert.InDelta(t, 0.0029, rows[0].CachedInputCostUSD, 1e-9, "legacy cache total lands on cached input")
assert.Zero(t, rows[0].CacheCreationCostUSD, "legacy rows carry no read/write split to recover")
assert.Zero(t, rows[0].OutputCostUSD, "legacy rows carry no input/output split to recover")
// u2: no cache spend — the whole total is the non-cache remainder.
assert.InDelta(t, 0.5, rows[1].TotalCostUSD(), 1e-9, "cache-free historical total must survive")
assert.Zero(t, rows[1].CacheCostUSD(), "a cache-free row must stay cache-free")
// u3: zero stays zero rather than being rewritten.
assert.Zero(t, rows[2].TotalCostUSD(), "an unpriced row must remain unpriced")
}
// TestFoldCostAggregatesIntoBuckets_SkipsAlreadyMigrated proves the migration is
// safe to re-run: with no legacy column present it is a no-op that leaves a
// true four-way split untouched.
func TestFoldCostAggregatesIntoBuckets_SkipsAlreadyMigrated(t *testing.T) {
ctx := context.Background()
db := setupDatabase(t)
require.NoError(t, db.Migrator().DropTable(&agentNetworkTypes.AgentNetworkUsage{}))
require.NoError(t, db.AutoMigrate(&agentNetworkTypes.AgentNetworkUsage{}))
// Timestamp must be set explicitly: a zero time.Time serialises as
// '0000-00-00 00:00:00', which MySQL rejects under strict mode.
require.NoError(t, db.Create(&agentNetworkTypes.AgentNetworkUsage{
ID: "u1", AccountID: "acct-1", Model: "claude-sonnet-4-6",
Timestamp: time.Date(2026, 5, 5, 9, 0, 0, 0, time.UTC),
InputCostUSD: 0.001, CachedInputCostUSD: 0.002, CacheCreationCostUSD: 0.003, OutputCostUSD: 0.004,
}).Error)
require.NoError(t, migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkUsage](ctx, db),
"running against an already-migrated table must be a no-op, not an error")
var row agentNetworkTypes.AgentNetworkUsage
require.NoError(t, db.First(&row, "id = ?", "u1").Error)
assert.InDelta(t, 0.001, row.InputCostUSD, 1e-9, "a true split must not be rewritten")
assert.InDelta(t, 0.002, row.CachedInputCostUSD, 1e-9)
assert.InDelta(t, 0.003, row.CacheCreationCostUSD, 1e-9)
assert.InDelta(t, 0.004, row.OutputCostUSD, 1e-9)
assert.InDelta(t, 0.01, row.TotalCostUSD(), 1e-9, "derived total sums the four buckets")
}

View File

@@ -14,7 +14,6 @@ import (
nbDomain "github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
type NetworkResourceType string
@@ -65,27 +64,6 @@ func NewNetworkResource(accountID, networkID, name, description, address string,
}, nil
}
// ToComponent converts the resource to its self-contained components
// representation. Returns nil for a nil resource.
func (n *NetworkResource) ToComponent() *sharedTypes.ComponentResource {
if n == nil {
return nil
}
return &sharedTypes.ComponentResource{
ID: n.ID,
PublicID: n.PublicID,
NetworkID: n.NetworkID,
AccountID: n.AccountID,
Name: n.Name,
Description: n.Description,
Type: sharedTypes.ComponentResourceType(n.Type),
Address: n.Address,
Domain: n.Domain,
Prefix: n.Prefix,
Enabled: n.Enabled,
}
}
func (n *NetworkResource) ToAPIResponse(groups []api.GroupMinimum) *api.NetworkResource {
addr := n.Prefix.String()
if n.Type == Domain {

View File

@@ -7,7 +7,6 @@ import (
"github.com/netbirdio/netbird/management/server/networks/types"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
type NetworkRouter struct {
@@ -22,36 +21,6 @@ type NetworkRouter struct {
Enabled bool
}
// ToComponent converts the router to its self-contained components
// representation. Returns nil for a nil router.
func (n *NetworkRouter) ToComponent() *sharedTypes.ComponentRouter {
if n == nil {
return nil
}
return &sharedTypes.ComponentRouter{
NetworkID: n.NetworkID,
PublicID: n.PublicID,
Peer: n.Peer,
PeerGroups: n.PeerGroups,
Masquerade: n.Masquerade,
Metric: n.Metric,
Enabled: n.Enabled,
}
}
// ToComponentMap converts a peer-keyed router map to its components
// representation.
func ToComponentMap(routers map[string]*NetworkRouter) map[string]*sharedTypes.ComponentRouter {
if routers == nil {
return nil
}
out := make(map[string]*sharedTypes.ComponentRouter, len(routers))
for id, r := range routers {
out[id] = r.ToComponent()
}
return out
}
func NewNetworkRouter(accountID string, networkID string, peer string, peerGroups []string, masquerade bool, metric int, enabled bool) (*NetworkRouter, error) {
r := &NetworkRouter{
ID: xid.New().String(),

View File

@@ -405,7 +405,7 @@ func (am *DefaultAccountManager) CreatePeerJob(ctx context.Context, accountID, p
return status.NewPeerNotPartOfAccountError()
}
meetMinVer, err := version.MeetsMinVersion(remoteJobsMinVer, p.Meta.WtVersion)
meetMinVer, err := posture.MeetsMinVersion(remoteJobsMinVer, p.Meta.WtVersion)
if !version.IsDevelopmentVersion(p.Meta.WtVersion) && (!meetMinVer || err != nil) {
return status.Errorf(status.PreconditionFailed, "peer version %s does not meet the minimum required version %s for remote jobs", p.Meta.WtVersion, remoteJobsMinVer)
}
@@ -1588,7 +1588,7 @@ func affectedPeerIDsFromNetworkMap(nmap *types.NetworkMap, selfPeerID string) []
}
seen := make(map[string]struct{}, len(nmap.Peers)+len(nmap.OfflinePeers))
ids := make([]string, 0, len(nmap.Peers)+len(nmap.OfflinePeers))
add := func(peers []*types.ComponentPeer) {
add := func(peers []*nbpeer.Peer) {
for _, p := range peers {
if p == nil || p.ID == "" || p.ID == selfPeerID {
continue

View File

@@ -13,7 +13,6 @@ import (
"github.com/netbirdio/netbird/management/server/util"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
// Peer capability constants mirror the proto enum values.
@@ -206,35 +205,6 @@ func (p *Peer) AddedWithSSOLogin() bool {
return p.UserID != ""
}
// ToComponent converts the peer to its self-contained components
// representation, carrying exactly the subset of peer data that crosses the
// components wire format. Returns nil for a nil peer so callers can convert
// possibly-missing peers without guarding.
func (p *Peer) ToComponent() *sharedTypes.ComponentPeer {
if p == nil {
return nil
}
cp := &sharedTypes.ComponentPeer{
ID: p.ID,
Key: p.Key,
IP: p.IP,
IPv6: p.IPv6,
DNSLabel: p.DNSLabel,
SSHKey: p.SSHKey,
SSHEnabled: p.SSHEnabled,
ServerSSHAllowed: p.Meta.Flags.ServerSSHAllowed,
AgentVersion: p.Meta.WtVersion,
SupportsSourcePrefixes: p.SupportsSourcePrefixes(),
SupportsIPv6: p.SupportsIPv6(),
LoginExpirationEnabled: p.LoginExpirationEnabled,
AddedWithSSOLogin: p.AddedWithSSOLogin(),
}
if p.LastLogin != nil {
cp.LastLogin = *p.LastLogin
}
return cp
}
// HasCapability reports whether the peer has the given capability.
func (p *Peer) HasCapability(capability int32) bool {
return slices.Contains(p.Meta.Capabilities, capability)

View File

@@ -1092,14 +1092,14 @@ func TestToSyncResponse(t *testing.T) {
}
networkMap := &types.NetworkMap{
Network: &types.Network{Net: *ipnet, Serial: 1000},
Peers: []*types.ComponentPeer{{
Peers: []*nbpeer.Peer{{
IP: netip.MustParseAddr("192.168.1.2"),
IPv6: netip.MustParseAddr("fd00::2"),
Key: "peer2-key",
DNSLabel: "peer2",
SSHEnabled: true,
SSHKey: "peer2-ssh-key"}},
OfflinePeers: []*types.ComponentPeer{{
OfflinePeers: []*nbpeer.Peer{{
IP: netip.MustParseAddr("192.168.1.3"),
IPv6: netip.MustParseAddr("fd00::3"),
Key: "peer3-key",

View File

@@ -3,9 +3,11 @@ package posture
import (
"context"
"fmt"
"strings"
"github.com/hashicorp/go-version"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbversion "github.com/netbirdio/netbird/version"
)
type NBVersionCheck struct {
@@ -14,8 +16,14 @@ type NBVersionCheck struct {
var _ Check = (*NBVersionCheck)(nil)
// sanitizeVersion removes anything after the pre-release tag (e.g., "-dev", "-alpha", etc.)
func sanitizeVersion(version string) string {
parts := strings.Split(version, "-")
return parts[0]
}
func (n *NBVersionCheck) Check(ctx context.Context, peer nbpeer.Peer) (bool, error) {
meetsMin, err := nbversion.MeetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
meetsMin, err := MeetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
if err != nil {
return false, err
}
@@ -40,3 +48,21 @@ func (n *NBVersionCheck) Validate() error {
}
return nil
}
// MeetsMinVersion checks if the peer's version meets or exceeds the minimum required version
func MeetsMinVersion(minVer, peerVer string) (bool, error) {
peerVer = sanitizeVersion(peerVer)
minVer = sanitizeVersion(minVer)
peerNBVer, err := version.NewVersion(peerVer)
if err != nil {
return false, err
}
constraints, err := version.NewConstraint(">= " + minVer)
if err != nil {
return false, err
}
return constraints.Check(peerNBVer), nil
}

View File

@@ -139,3 +139,68 @@ func TestNBVersionCheck_Validate(t *testing.T) {
})
}
}
func TestMeetsMinVersion(t *testing.T) {
tests := []struct {
name string
minVer string
peerVer string
want bool
wantErr bool
}{
{
name: "Peer version greater than min version",
minVer: "0.26.0",
peerVer: "0.60.1",
want: true,
wantErr: false,
},
{
name: "Peer version equals min version",
minVer: "1.0.0",
peerVer: "1.0.0",
want: true,
wantErr: false,
},
{
name: "Peer version less than min version",
minVer: "1.0.0",
peerVer: "0.9.9",
want: false,
wantErr: false,
},
{
name: "Peer version with pre-release tag greater than min version",
minVer: "1.0.0",
peerVer: "1.0.1-alpha",
want: true,
wantErr: false,
},
{
name: "Invalid peer version format",
minVer: "1.0.0",
peerVer: "dev",
want: false,
wantErr: true,
},
{
name: "Invalid min version format",
minVer: "invalid.version",
peerVer: "1.0.0",
want: false,
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := MeetsMinVersion(tt.minVer, tt.peerVer)
if tt.wantErr {
assert.Error(t, err)
} else {
assert.NoError(t, err)
}
assert.Equal(t, tt.want, got)
})
}
}

View File

@@ -71,7 +71,7 @@ func (s *SqlStore) GetAgentNetworkMetrics(ctx context.Context) (AgentNetworkMetr
usageRow := db.Model(&agentNetworkTypes.AgentNetworkUsage{}).
Select("COALESCE(SUM(input_tokens), 0) AS input_tokens, " +
"COALESCE(SUM(output_tokens), 0) AS output_tokens, " +
"COALESCE(SUM" + agentNetworkTypes.CostUSDSQLExpr + ", 0) AS cost_usd").Row()
"COALESCE(SUM(cost_usd), 0) AS cost_usd").Row()
if err := usageRow.Scan(&m.InputTokens, &m.OutputTokens, &m.CostUSD); err != nil {
return AgentNetworkMetrics{}, fmt.Errorf("scan agent network usage metrics: %w", err)
}

View File

@@ -37,7 +37,7 @@ func TestAgentNetworkUsage_RealStore_RoundTrip(t *testing.T) {
InputTokens: 1200,
OutputTokens: 640,
TotalTokens: 1840,
InputCostUSD: 0.0231,
CostUSD: 0.0231,
}
usageGroups := []agentNetworkTypes.AgentNetworkUsageGroup{
{UsageID: usage.ID, GroupID: "grp-eng", AccountID: accountID},
@@ -71,7 +71,7 @@ func TestAgentNetworkUsage_RealStore_RoundTrip(t *testing.T) {
InputTokens: 1200,
OutputTokens: 640,
TotalTokens: 1840,
InputCostUSD: 0.0231,
CostUSD: 0.0231,
}
entryGroups := []agentNetworkTypes.AgentNetworkAccessLogGroup{
{LogID: entry.ID, GroupID: "grp-eng", AccountID: accountID},
@@ -127,7 +127,7 @@ func TestAgentNetworkUsageOverview_DailyAggregation(t *testing.T) {
mk := func(id string, ts time.Time, model string, in, out int64, cost float64) *agentNetworkTypes.AgentNetworkUsage {
return &agentNetworkTypes.AgentNetworkUsage{
ID: id, AccountID: accountID, Timestamp: ts, Model: model,
InputTokens: in, OutputTokens: out, TotalTokens: in + out, InputCostUSD: cost,
InputTokens: in, OutputTokens: out, TotalTokens: in + out, CostUSD: cost,
}
}
require.NoError(t, s.CreateAgentNetworkUsage(ctx, mk("u1", day1, "gpt-4o", 100, 50, 0.10), nil))
@@ -143,7 +143,7 @@ func TestAgentNetworkUsageOverview_DailyAggregation(t *testing.T) {
assert.Equal(t, "2026-05-05", buckets[0].PeriodStart, "oldest-first ordering")
assert.Equal(t, int64(300), buckets[0].InputTokens, "same-day input tokens summed")
assert.Equal(t, int64(130), buckets[0].OutputTokens)
assert.InDelta(t, 0.30, buckets[0].TotalCostUSD(), 1e-9, "same-day cost summed")
assert.InDelta(t, 0.30, buckets[0].CostUSD, 1e-9, "same-day cost summed")
assert.Equal(t, "2026-05-06", buckets[1].PeriodStart)
assert.Equal(t, int64(15), buckets[1].TotalTokens)
@@ -174,7 +174,7 @@ func TestAgentNetworkAccessLogSessions_RealStore(t *testing.T) {
ID: id, AccountID: accountID, ServiceID: "svc", Timestamp: ts,
UserID: user, StatusCode: 200, Provider: provider, Model: model,
SessionID: session, Decision: decision,
InputTokens: 100, OutputTokens: 50, TotalTokens: 150, InputCostUSD: cost,
InputTokens: 100, OutputTokens: 50, TotalTokens: 150, CostUSD: cost,
}
}
@@ -207,7 +207,7 @@ func TestAgentNetworkAccessLogSessions_RealStore(t *testing.T) {
s1 := sessions[2]
assert.Equal(t, 2, s1.RequestCount, "s1 has two requests")
assert.Equal(t, int64(300), s1.TotalTokens, "tokens summed across the session")
assert.InDelta(t, 0.30, s1.TotalCostUSD(), 1e-9, "cost summed across the session")
assert.InDelta(t, 0.30, s1.CostUSD, 1e-9, "cost summed across the session")
assert.Equal(t, "alice", s1.UserID)
assert.Equal(t, "allow", s1.Decision)
// SQLite hands times back in time.Local; normalise to UTC so the instant is

View File

@@ -650,14 +650,6 @@ func getMigrationsPostAuto(ctx context.Context) []migrationFunc {
func(db *gorm.DB) error {
return migration.DropIndex[proxy.Proxy](ctx, db, "idx_proxy_account_id_unique")
},
// Post-auto so the per-bucket cost columns already exist when the legacy
// aggregates are folded into them and dropped.
func(db *gorm.DB) error {
return migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkAccessLog](ctx, db)
},
func(db *gorm.DB) error {
return migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkUsage](ctx, db)
},
}
}

View File

@@ -283,8 +283,8 @@ func (a *Account) SynthesizePrivateServiceZones(peerID string) []nbdns.CustomZon
// it, adding a single private service would black-hole every
// other name under the zone apex.
zone = &nbdns.CustomZone{
Domain: dns.Fqdn(serviceDomainZone),
Records: []nbdns.SimpleRecord{},
Domain: dns.Fqdn(serviceDomainZone),
Records: []nbdns.SimpleRecord{},
NonAuthoritative: true,
SearchDomainDisabled: true,
}
@@ -1082,7 +1082,6 @@ func (a *Account) connResourcesGenerator(ctx context.Context, targetPeer *nbpeer
peersExists := make(map[string]struct{})
rules := make([]*FirewallRule, 0)
peers := make([]*nbpeer.Peer, 0)
targetComponent := targetPeer.ToComponent()
return func(rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int) {
for _, peer := range groupPeers {
@@ -1118,10 +1117,10 @@ func (a *Account) connResourcesGenerator(ctx context.Context, targetPeer *nbpeer
if len(rule.Ports) == 0 && len(rule.PortRanges) == 0 {
rules = append(rules, &fr)
} else {
rules = append(rules, ExpandPortsAndRanges(fr, rule, targetComponent)...)
rules = append(rules, ExpandPortsAndRanges(fr, rule, targetPeer)...)
}
rules = AppendIPv6FirewallRule(rules, rulesExists, peer.ToComponent(), targetComponent, rule, FirewallRuleContext{
rules = AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, FirewallRuleContext{
Direction: direction,
DirStr: strconv.Itoa(direction),
ProtocolStr: string(protocol),
@@ -1281,7 +1280,7 @@ func (a *Account) getRouteFirewallRules(ctx context.Context, peerID string, poli
return fwRules
}
func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}, validatedPeersMap map[string]struct{}) []*ComponentPeer {
func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}, validatedPeersMap map[string]struct{}) []*nbpeer.Peer {
distPeersWithPolicy := make(map[string]struct{})
for _, id := range rule.Sources {
group := a.Groups[id]
@@ -1308,13 +1307,13 @@ func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID
}
}
distributionGroupPeers := make([]*ComponentPeer, 0, len(distPeersWithPolicy))
distributionGroupPeers := make([]*nbpeer.Peer, 0, len(distPeersWithPolicy))
for pID := range distPeersWithPolicy {
peer := a.Peers[pID]
if peer == nil {
continue
}
distributionGroupPeers = append(distributionGroupPeers, peer.ToComponent())
distributionGroupPeers = append(distributionGroupPeers, peer)
}
return distributionGroupPeers
}

View File

@@ -9,7 +9,9 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/internals/modules/zones"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/route"
)
@@ -111,7 +113,7 @@ func (a *Account) GetPeerNetworkMapComponents(
PeerID: peerID,
Network: a.Network.Copy(),
// must include the target peer as it's required on the client
Peers: map[string]*ComponentPeer{peerID: peer.ToComponent()},
Peers: map[string]*nbpeer.Peer{peerID: peer},
})
}
@@ -124,7 +126,7 @@ func (a *Account) GetPeerNetworkMapComponents(
PeerID: peerID,
Network: a.Network.Copy(),
// must include the target peer as it's required on the client
Peers: map[string]*ComponentPeer{peerID: peer.ToComponent()},
Peers: map[string]*nbpeer.Peer{peerID: peer},
})
}
@@ -134,10 +136,10 @@ func (a *Account) GetPeerNetworkMapComponents(
NameServerGroups: make([]*nbdns.NameServerGroup, 0),
CustomZoneDomain: peersCustomZone.Domain,
ResourcePoliciesMap: make(map[string][]*Policy),
RoutersMap: make(map[string]map[string]*ComponentRouter),
NetworkResources: make([]*ComponentResource, 0),
RoutersMap: make(map[string]map[string]*routerTypes.NetworkRouter),
NetworkResources: make([]*resourceTypes.NetworkResource, 0),
PostureFailedPeers: make(map[string]map[string]struct{}, len(a.PostureChecks)),
RouterPeers: make(map[string]*ComponentPeer),
RouterPeers: make(map[string]*nbpeer.Peer),
NetworkXIDToPublicID: make(map[string]string, len(a.Networks)),
PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)),
}
@@ -172,7 +174,7 @@ func (a *Account) GetPeerNetworkMapComponents(
}
components.Peers = relevantPeers
components.Groups = GroupsToComponent(relevantGroups)
components.Groups = relevantGroups
components.Policies = relevantPolicies
components.Routes = relevantRoutes
components.AllDNSRecords = filterDNSRecordsByPeers(peersCustomZone.Records, relevantPeers, peer.SupportsIPv6() && peer.IPv6.IsValid())
@@ -221,7 +223,7 @@ func (a *Account) GetPeerNetworkMapComponents(
}
for _, pID := range a.getPostureValidPeersSaveFailed(peers, policy.SourcePostureChecks, validatedPeersMap, &components.PostureFailedPeers) {
if _, exists := components.Peers[pID]; !exists {
components.Peers[pID] = a.GetPeer(pID).ToComponent()
components.Peers[pID] = a.GetPeer(pID)
}
}
} else {
@@ -254,14 +256,14 @@ func (a *Account) GetPeerNetworkMapComponents(
for _, srcGroupID := range rule.Sources {
if g := a.Groups[srcGroupID]; g != nil {
if _, exists := components.Groups[srcGroupID]; !exists {
components.Groups[srcGroupID] = g.ToComponent()
components.Groups[srcGroupID] = g
}
}
}
for _, dstGroupID := range rule.Destinations {
if g := a.Groups[dstGroupID]; g != nil {
if _, exists := components.Groups[dstGroupID]; !exists {
components.Groups[dstGroupID] = g.ToComponent()
components.Groups[dstGroupID] = g
}
}
}
@@ -276,22 +278,20 @@ func (a *Account) GetPeerNetworkMapComponents(
// network in the account — accounts with many tenants/networks
// shipped tens of unrelated peers in `peers[]` and `routers_map`.
if addSourcePeers {
components.RoutersMap[resource.NetworkID] = routerTypes.ToComponentMap(networkRoutingPeers)
components.RoutersMap[resource.NetworkID] = networkRoutingPeers
for peerIDKey := range networkRoutingPeers {
if p := a.Peers[peerIDKey]; p != nil {
cp := components.RouterPeers[peerIDKey]
if cp == nil {
cp = p.ToComponent()
components.RouterPeers[peerIDKey] = cp
if _, exists := components.RouterPeers[peerIDKey]; !exists {
components.RouterPeers[peerIDKey] = p
}
if _, exists := components.Peers[peerIDKey]; !exists {
if _, validated := validatedPeersMap[peerIDKey]; validated {
components.Peers[peerIDKey] = cp
components.Peers[peerIDKey] = p
}
}
}
}
components.NetworkResources = append(components.NetworkResources, resource.ToComponent())
components.NetworkResources = append(components.NetworkResources, resource)
}
}
@@ -312,14 +312,14 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
peerSSHEnabled bool,
validatedPeersMap map[string]struct{},
postureFailedPeers *map[string]map[string]struct{},
) (map[string]*ComponentPeer, map[string]*Group, []*Policy, []*route.Route, sshRequirements) {
relevantPeerIDs := make(map[string]*ComponentPeer, len(a.Peers)/4)
) (map[string]*nbpeer.Peer, map[string]*Group, []*Policy, []*route.Route, sshRequirements) {
relevantPeerIDs := make(map[string]*nbpeer.Peer, len(a.Peers)/4)
relevantGroupIDs := make(map[string]*Group, len(a.Groups)/4)
relevantPolicies := make([]*Policy, 0, len(a.Policies))
relevantRoutes := make([]*route.Route, 0, len(a.Routes))
sshReqs := sshRequirements{neededGroupIDs: make(map[string]struct{})}
relevantPeerIDs[peerID] = a.GetPeer(peerID).ToComponent()
relevantPeerIDs[peerID] = a.GetPeer(peerID)
peerGroupSet := make(map[string]struct{}, 8)
for groupID, group := range a.Groups {
@@ -384,7 +384,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
if r.Peer != "" {
if _, ok := validatedPeersMap[r.Peer]; ok {
if p := a.GetPeer(r.Peer); p != nil {
relevantPeerIDs[r.Peer] = p.ToComponent()
relevantPeerIDs[r.Peer] = p
}
}
}
@@ -401,7 +401,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
continue
}
if p := a.GetPeer(pid); p != nil {
relevantPeerIDs[pid] = p.ToComponent()
relevantPeerIDs[pid] = p
}
}
}
@@ -458,9 +458,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
if peerInSources {
policyRelevant = true
for _, pid := range destinationPeers {
if _, exists := relevantPeerIDs[pid]; !exists {
relevantPeerIDs[pid] = a.GetPeer(pid).ToComponent()
}
relevantPeerIDs[pid] = a.GetPeer(pid)
}
for _, dstGroupID := range rule.Destinations {
relevantGroupIDs[dstGroupID] = a.GetGroup(dstGroupID)
@@ -470,9 +468,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
if peerInDestinations {
policyRelevant = true
for _, pid := range sourcePeers {
if _, exists := relevantPeerIDs[pid]; !exists {
relevantPeerIDs[pid] = a.GetPeer(pid).ToComponent()
}
relevantPeerIDs[pid] = a.GetPeer(pid)
}
for _, srcGroupID := range rule.Sources {
relevantGroupIDs[srcGroupID] = a.GetGroup(srcGroupID)
@@ -628,7 +624,7 @@ func (a *Account) getPostureValidPeersSaveFailed(inputPeers []string, postureChe
// that name them. Calculate() tolerates groups with empty Peers (the inner
// loops simply iterate zero times), so retaining them is behaviourally a
// no-op for the legacy path that consumes the same NetworkMapComponents.
func filterGroupPeers(groups *map[string]*ComponentGroup, peers map[string]*ComponentPeer) {
func filterGroupPeers(groups *map[string]*Group, peers map[string]*nbpeer.Peer) {
for groupID, groupInfo := range *groups {
filteredPeers := make([]string, 0, len(groupInfo.Peers))
for _, pid := range groupInfo.Peers {
@@ -638,14 +634,14 @@ func filterGroupPeers(groups *map[string]*ComponentGroup, peers map[string]*Comp
}
if len(filteredPeers) != len(groupInfo.Peers) {
ng := *groupInfo
ng := groupInfo.Copy()
ng.Peers = filteredPeers
(*groups)[groupID] = &ng
(*groups)[groupID] = ng
}
}
}
func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}, policies []*Policy, resourcePoliciesMap map[string][]*Policy, peers map[string]*ComponentPeer) {
func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}, policies []*Policy, resourcePoliciesMap map[string][]*Policy, peers map[string]*nbpeer.Peer) {
if len(*postureFailedPeers) == 0 {
return
}
@@ -680,7 +676,7 @@ func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}
}
}
func filterDNSRecordsByPeers(records []nbdns.SimpleRecord, peers map[string]*ComponentPeer, includeIPv6 bool) []nbdns.SimpleRecord {
func filterDNSRecordsByPeers(records []nbdns.SimpleRecord, peers map[string]*nbpeer.Peer, includeIPv6 bool) []nbdns.SimpleRecord {
if len(records) == 0 || len(peers) == 0 {
return nil
}

View File

@@ -9,6 +9,7 @@ import (
"github.com/stretchr/testify/require"
nbdns "github.com/netbirdio/netbird/dns"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
)
func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
@@ -48,7 +49,7 @@ func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
})
}
func netmapPeerIDs(peers []*ComponentPeer) []string {
func netmapPeerIDs(peers []*nbpeer.Peer) []string {
ids := make([]string, 0, len(peers))
for _, p := range peers {
if p == nil {

View File

@@ -666,7 +666,7 @@ func Test_ExpandPortsAndRanges_SSHRuleExpansion(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := ExpandPortsAndRanges(tt.base, tt.rule, tt.peer.ToComponent())
result := ExpandPortsAndRanges(tt.base, tt.rule, tt.peer)
var ports []string
for _, fr := range result {

View File

@@ -6,6 +6,7 @@ import (
"net"
"net/netip"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbroute "github.com/netbirdio/netbird/route"
sharedtypes "github.com/netbirdio/netbird/shared/management/types"
)
@@ -17,6 +18,9 @@ type DNSSettings = sharedtypes.DNSSettings
type FirewallRule = sharedtypes.FirewallRule
type Group = sharedtypes.Group
type GroupPeer = sharedtypes.GroupPeer
type Network = sharedtypes.Network
type NetworkMap = sharedtypes.NetworkMap
type ForwardingRule = sharedtypes.ForwardingRule
@@ -38,18 +42,6 @@ type RouteFirewallRule = sharedtypes.RouteFirewallRule
type NetworkMapComponents = sharedtypes.NetworkMapComponents
type ComponentPeer = sharedtypes.ComponentPeer
type ComponentGroup = sharedtypes.ComponentGroup
type ComponentRouter = sharedtypes.ComponentRouter
type ComponentResource = sharedtypes.ComponentResource
type ComponentResourceType = sharedtypes.ComponentResourceType
const (
ComponentResourceHost = sharedtypes.ComponentResourceHost
ComponentResourceSubnet = sharedtypes.ComponentResourceSubnet
ComponentResourceDomain = sharedtypes.ComponentResourceDomain
)
var EmptyNetworkMapComponents = sharedtypes.EmptyNetworkMapComponents
type AccountSettingsInfo = sharedtypes.AccountSettingsInfo
@@ -60,7 +52,12 @@ type NetworkMapComponentsCompact = sharedtypes.NetworkMapComponentsCompact
type LookupMap = sharedtypes.LookupMap
type FirewallRuleContext = sharedtypes.FirewallRuleContext
const GroupAllName = sharedtypes.GroupAllName
const (
GroupIssuedAPI = sharedtypes.GroupIssuedAPI
GroupIssuedJWT = sharedtypes.GroupIssuedJWT
GroupIssuedIntegration = sharedtypes.GroupIssuedIntegration
GroupAllName = sharedtypes.GroupAllName
)
// Function forwarders preserve types.X(...) call sites that previously
// resolved to package-local funcs. Plain forwarders (not var aliases) keep
@@ -70,11 +67,11 @@ func PolicyRuleImpliesLegacySSH(rule *PolicyRule) bool {
return sharedtypes.PolicyRuleImpliesLegacySSH(rule)
}
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *ComponentPeer) []*FirewallRule {
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *nbpeer.Peer) []*FirewallRule {
return sharedtypes.ExpandPortsAndRanges(base, rule, peer)
}
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *ComponentPeer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *nbpeer.Peer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
return sharedtypes.AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, rc)
}
@@ -82,7 +79,7 @@ func CalculateNetworkMapFromComponents(ctx context.Context, components *NetworkM
return sharedtypes.CalculateNetworkMapFromComponents(ctx, components)
}
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*ComponentPeer, direction int, includeIPv6 bool) []*RouteFirewallRule {
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int, includeIPv6 bool) []*RouteFirewallRule {
return sharedtypes.GenerateRouteFirewallRules(ctx, route, rule, groupPeers, direction, includeIPv6)
}

View File

@@ -9,7 +9,6 @@ import (
"github.com/stretchr/testify/require"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
)
func TestNetworkMapComponents_IPv6EndToEnd(t *testing.T) {
@@ -105,7 +104,7 @@ func TestNetworkMapComponents_RemotePeerWithoutCapability(t *testing.T) {
require.NotNil(t, nm)
t.Run("AllowedIPs include remote v6", func(t *testing.T) {
var dst *types.ComponentPeer
var dst *nbpeer.Peer
for _, p := range nm.Peers {
if p.ID == "peer-dst-1" {
dst = p

View File

@@ -49,7 +49,7 @@ func allPeersValidated(account *types.Account, excludePeerIDs ...string) map[str
return validated
}
func peerIDs(peers []*types.ComponentPeer) []string {
func peerIDs(peers []*nbpeer.Peer) []string {
ids := make([]string, len(peers))
for i, p := range peers {
ids[i] = p.ID

View File

@@ -19,3 +19,34 @@ func Difference(a, b []string) []string {
func ToPtr[T any](value T) *T {
return &value
}
type comparableObject[T any] interface {
Equal(other T) bool
}
func MergeUnique[T comparableObject[T]](arr1, arr2 []T) []T {
var result []T
for _, item := range arr1 {
if !contains(result, item) {
result = append(result, item)
}
}
for _, item := range arr2 {
if !contains(result, item) {
result = append(result, item)
}
}
return result
}
func contains[T comparableObject[T]](slice []T, element T) bool {
for _, item := range slice {
if item.Equal(element) {
return true
}
}
return false
}

Some files were not shown because too many files have changed in this diff Show More