mirror of
https://github.com/netbirdio/netbird.git
synced 2026-07-30 04:11:28 +02:00
Compare commits
17 Commits
feature/ch
...
agent-netw
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e6a286ab67 | ||
|
|
2316c20a2b | ||
|
|
dd2bdc0de3 | ||
|
|
44fef45c2f | ||
|
|
3d1f209ea3 | ||
|
|
2ef457be95 | ||
|
|
63c320b6a9 | ||
|
|
4acbe2670a | ||
|
|
0fb4c8c423 | ||
|
|
42e45ff9f9 | ||
|
|
9269b56386 | ||
|
|
b3f9b82442 | ||
|
|
8a43f4f943 | ||
|
|
2f268c8141 | ||
|
|
e3c4128164 | ||
|
|
bab5572a74 | ||
|
|
9b4a5df925 |
@@ -24,6 +24,8 @@ builds:
|
||||
ldflags:
|
||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||
tags:
|
||||
- production
|
||||
|
||||
- id: netbird-ui-windows-amd64
|
||||
dir: client/ui
|
||||
@@ -39,6 +41,8 @@ builds:
|
||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||
- -H windowsgui
|
||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||
tags:
|
||||
- production
|
||||
|
||||
- id: netbird-ui-windows-arm64
|
||||
dir: client/ui
|
||||
@@ -55,6 +59,8 @@ builds:
|
||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||
- -H windowsgui
|
||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||
tags:
|
||||
- production
|
||||
|
||||
archives:
|
||||
- id: linux-arch
|
||||
|
||||
@@ -29,6 +29,8 @@ builds:
|
||||
ldflags:
|
||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||
tags:
|
||||
- production
|
||||
|
||||
universal_binaries:
|
||||
- id: netbird-ui-darwin
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -299,6 +300,13 @@ func (c *Client) SetInfoLogLevel() {
|
||||
// PeersList return with the list of the PeerInfos
|
||||
func (c *Client) PeersList() *PeerInfoArray {
|
||||
|
||||
// The recorder only caches transfer counters and handshake times; nothing
|
||||
// refreshes them on its own, so without this they read as zero. The desktop
|
||||
// daemon does the same before serving a full peer status.
|
||||
if err := c.recorder.RefreshWireGuardStats(); err != nil {
|
||||
log.Debugf("failed to refresh WireGuard stats: %v", err)
|
||||
}
|
||||
|
||||
fullStatus := c.recorder.GetFullStatus()
|
||||
|
||||
peerInfos := make([]PeerInfo, len(fullStatus.Peers))
|
||||
@@ -309,6 +317,20 @@ func (c *Client) PeersList() *PeerInfoArray {
|
||||
FQDN: p.FQDN,
|
||||
ConnStatus: int(p.ConnStatus),
|
||||
Routes: PeerRoutes{routes: maps.Keys(p.GetRoutes())},
|
||||
|
||||
PubKey: p.PubKey,
|
||||
Latency: formatDuration(p.Latency),
|
||||
LatencyMs: p.Latency.Milliseconds(),
|
||||
BytesRx: p.BytesRx,
|
||||
BytesTx: p.BytesTx,
|
||||
ConnStatusUpdate: formatTime(p.ConnStatusUpdate),
|
||||
Relayed: p.Relayed,
|
||||
RosenpassEnabled: p.RosenpassEnabled,
|
||||
LastWireguardHandshake: formatTime(p.LastWireguardHandshake),
|
||||
LocalIceCandidateType: p.LocalIceCandidateType,
|
||||
RemoteIceCandidateType: p.RemoteIceCandidateType,
|
||||
LocalIceCandidateEndpoint: p.LocalIceCandidateEndpoint,
|
||||
RemoteIceCandidateEndpoint: p.RemoteIceCandidateEndpoint,
|
||||
}
|
||||
peerInfos[n] = pi
|
||||
}
|
||||
@@ -439,10 +461,6 @@ func (c *Client) RemoveConnectionListener() {
|
||||
c.recorder.RemoveConnectionListener()
|
||||
}
|
||||
|
||||
func (c *Client) toggleRoute(command routeCommand) error {
|
||||
return command.toggleRoute()
|
||||
}
|
||||
|
||||
func (c *Client) getRouteManager() (routemanager.Manager, error) {
|
||||
client := c.getConnectClient()
|
||||
if client == nil {
|
||||
@@ -462,22 +480,22 @@ func (c *Client) getRouteManager() (routemanager.Manager, error) {
|
||||
return manager, nil
|
||||
}
|
||||
|
||||
func (c *Client) SelectRoute(route string) error {
|
||||
func (c *Client) SelectRoute(id string) error {
|
||||
manager, err := c.getRouteManager()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return c.toggleRoute(selectRouteCommand{route: route, manager: manager})
|
||||
return manager.SelectRoutes([]route.NetID{route.NetID(id)}, true)
|
||||
}
|
||||
|
||||
func (c *Client) DeselectRoute(route string) error {
|
||||
func (c *Client) DeselectRoute(id string) error {
|
||||
manager, err := c.getRouteManager()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return c.toggleRoute(deselectRouteCommand{route: route, manager: manager})
|
||||
return manager.DeselectRoutes([]route.NetID{route.NetID(id)})
|
||||
}
|
||||
|
||||
// getNetworkDomainsFromRoute extracts domains from a route and enriches each domain
|
||||
@@ -512,3 +530,28 @@ func exportEnvList(list *EnvList) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// formatDuration renders a duration for display, trimming the fractional part
|
||||
// to two digits so latencies read as "12.34ms" rather than "12.345678ms".
|
||||
func formatDuration(d time.Duration) string {
|
||||
ds := d.String()
|
||||
dotIndex := strings.Index(ds, ".")
|
||||
if dotIndex == -1 {
|
||||
return ds
|
||||
}
|
||||
|
||||
endIndex := min(dotIndex+3, len(ds))
|
||||
|
||||
// Skip the remaining digits so only the unit suffix is appended back.
|
||||
unitStart := endIndex
|
||||
for unitStart < len(ds) && ds[unitStart] >= '0' && ds[unitStart] <= '9' {
|
||||
unitStart++
|
||||
}
|
||||
return ds[:endIndex] + ds[unitStart:]
|
||||
}
|
||||
|
||||
// formatTime renders a timestamp in UTC using a fixed layout. The zero time is
|
||||
// passed through as-is so the UI can recognise it and show "never" instead.
|
||||
func formatTime(t time.Time) string {
|
||||
return t.UTC().Format("2006-01-02 15:04:05")
|
||||
}
|
||||
|
||||
@@ -12,12 +12,30 @@ const (
|
||||
)
|
||||
|
||||
// PeerInfo describe information about the peers. It designed for the UI usage
|
||||
//
|
||||
// The fields below ConnStatus back the peer detail screen. Durations and times
|
||||
// are pre-formatted into strings so the UI does not have to know Go's layouts;
|
||||
// Latency is additionally exposed as LatencyMs for colour coding.
|
||||
type PeerInfo struct {
|
||||
IP string
|
||||
IPv6 string
|
||||
FQDN string
|
||||
ConnStatus int
|
||||
Routes PeerRoutes
|
||||
|
||||
PubKey string
|
||||
Latency string
|
||||
LatencyMs int64
|
||||
BytesRx int64
|
||||
BytesTx int64
|
||||
ConnStatusUpdate string
|
||||
Relayed bool
|
||||
RosenpassEnabled bool
|
||||
LastWireguardHandshake string
|
||||
LocalIceCandidateType string
|
||||
RemoteIceCandidateType string
|
||||
LocalIceCandidateEndpoint string
|
||||
RemoteIceCandidateEndpoint string
|
||||
}
|
||||
|
||||
func (p *PeerInfo) GetPeerRoutes() *PeerRoutes {
|
||||
|
||||
@@ -189,6 +189,19 @@ func (pm *ProfileManager) LogoutProfile(id string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RenameProfile changes a profile's display name. The profile ID, and therefore
|
||||
// its on-disk filename, is left untouched: only the "name" field of the config
|
||||
// is rewritten. This works for the default profile too, whose config lives in
|
||||
// netbird.cfg rather than under profiles/.
|
||||
func (pm *ProfileManager) RenameProfile(id string, newName string) error {
|
||||
if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), androidUsername, newName); err != nil {
|
||||
return fmt.Errorf("failed to rename profile: %w", err)
|
||||
}
|
||||
|
||||
log.Infof("renamed profile %s to: %s", id, newName)
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveProfile deletes a profile
|
||||
func (pm *ProfileManager) RemoveProfile(id string) error {
|
||||
// Use ServiceManager (removes profile from profiles/ directory)
|
||||
|
||||
@@ -1,70 +0,0 @@
|
||||
//go:build android
|
||||
|
||||
package android
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/routemanager"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
func executeRouteToggle(id string, manager routemanager.Manager,
|
||||
operationName string,
|
||||
routeOperation func(routes []route.NetID, allRoutes []route.NetID) error) error {
|
||||
netID := route.NetID(id)
|
||||
routes := []route.NetID{netID}
|
||||
|
||||
routesMap := manager.GetClientRoutesWithNetID()
|
||||
routes = route.ExpandV6ExitPairs(routes, routesMap)
|
||||
|
||||
log.Debugf("%s with ids: %v", operationName, routes)
|
||||
|
||||
if err := routeOperation(routes, maps.Keys(routesMap)); err != nil {
|
||||
log.Debugf("error when %s: %s", operationName, err)
|
||||
return fmt.Errorf("error %s: %w", operationName, err)
|
||||
}
|
||||
|
||||
manager.TriggerSelection(manager.GetClientRoutes())
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type routeCommand interface {
|
||||
toggleRoute() error
|
||||
}
|
||||
|
||||
type selectRouteCommand struct {
|
||||
route string
|
||||
manager routemanager.Manager
|
||||
}
|
||||
|
||||
func (s selectRouteCommand) toggleRoute() error {
|
||||
routeSelector := s.manager.GetRouteSelector()
|
||||
if routeSelector == nil {
|
||||
return fmt.Errorf("no route selector available")
|
||||
}
|
||||
|
||||
routeOperation := func(routes []route.NetID, allRoutes []route.NetID) error {
|
||||
return routeSelector.SelectRoutes(routes, true, allRoutes)
|
||||
}
|
||||
|
||||
return executeRouteToggle(s.route, s.manager, "selecting route", routeOperation)
|
||||
}
|
||||
|
||||
type deselectRouteCommand struct {
|
||||
route string
|
||||
manager routemanager.Manager
|
||||
}
|
||||
|
||||
func (d deselectRouteCommand) toggleRoute() error {
|
||||
routeSelector := d.manager.GetRouteSelector()
|
||||
if routeSelector == nil {
|
||||
return fmt.Errorf("no route selector available")
|
||||
}
|
||||
|
||||
return executeRouteToggle(d.route, d.manager, "deselecting route", routeSelector.DeselectRoutes)
|
||||
}
|
||||
@@ -385,11 +385,20 @@ func inactivityThresholdEnv() *time.Duration {
|
||||
return nil
|
||||
}
|
||||
|
||||
parsedMinutes, err := strconv.Atoi(envValue)
|
||||
if err != nil || parsedMinutes <= 0 {
|
||||
return nil
|
||||
// Documented format: a Go duration such as "30m" or "1h".
|
||||
if d, err := time.ParseDuration(envValue); err == nil {
|
||||
if d <= 0 {
|
||||
return nil
|
||||
}
|
||||
return &d
|
||||
}
|
||||
|
||||
d := time.Duration(parsedMinutes) * time.Minute
|
||||
return &d
|
||||
// Backwards compatibility: a bare integer used to be interpreted as minutes.
|
||||
if parsedMinutes, err := strconv.Atoi(envValue); err == nil && parsedMinutes > 0 {
|
||||
d := time.Duration(parsedMinutes) * time.Minute
|
||||
return &d
|
||||
}
|
||||
|
||||
log.Warnf("invalid %s value %q: expected a Go duration such as 30m or 1h", lazyconn.EnvInactivityThreshold, envValue)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -104,3 +104,38 @@ func TestConnMgr_ActivatePeerConcurrentWithLifecycle(t *testing.T) {
|
||||
close(done)
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestInactivityThresholdEnv(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
val string
|
||||
want *time.Duration
|
||||
}{
|
||||
{name: "unset", val: "", want: nil},
|
||||
{name: "go duration minutes", val: "30m", want: durPtr(30 * time.Minute)},
|
||||
{name: "go duration hours", val: "1h", want: durPtr(time.Hour)},
|
||||
{name: "go duration seconds", val: "90s", want: durPtr(90 * time.Second)},
|
||||
{name: "bare integer is minutes (backwards compat)", val: "5", want: durPtr(5 * time.Minute)},
|
||||
{name: "zero duration", val: "0s", want: nil},
|
||||
{name: "zero integer", val: "0", want: nil},
|
||||
{name: "negative duration", val: "-5m", want: nil},
|
||||
{name: "garbage", val: "abc", want: nil},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Setenv(lazyconn.EnvInactivityThreshold, tc.val)
|
||||
got := inactivityThresholdEnv()
|
||||
switch {
|
||||
case tc.want == nil && got != nil:
|
||||
t.Fatalf("want nil, got %v", *got)
|
||||
case tc.want != nil && got == nil:
|
||||
t.Fatalf("want %v, got nil", *tc.want)
|
||||
case tc.want != nil && *got != *tc.want:
|
||||
t.Fatalf("want %v, got %v", *tc.want, *got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func durPtr(d time.Duration) *time.Duration { return &d }
|
||||
|
||||
@@ -34,6 +34,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
"github.com/netbirdio/netbird/client/internal/statemanager"
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
"github.com/netbirdio/netbird/client/internal/tunnelnotifier"
|
||||
"github.com/netbirdio/netbird/client/internal/updater"
|
||||
"github.com/netbirdio/netbird/client/internal/updater/installer"
|
||||
nbnet "github.com/netbirdio/netbird/client/net"
|
||||
@@ -136,10 +137,13 @@ func (c *ConnectClient) RunOniOS(
|
||||
// Set GC percent to 5% to reduce memory usage as iOS only allows 50MB of memory for the extension.
|
||||
debug.SetGCPercent(5)
|
||||
|
||||
notifier := tunnelnotifier.New(networkChangeListener, dnsManager)
|
||||
defer notifier.Close()
|
||||
|
||||
mobileDependency := MobileDependency{
|
||||
FileDescriptor: fileDescriptor,
|
||||
NetworkChangeListener: networkChangeListener,
|
||||
DnsManager: dnsManager,
|
||||
NetworkChangeListener: notifier,
|
||||
DnsManager: notifier,
|
||||
StateFilePath: stateFilePath,
|
||||
TempDir: cacheDir,
|
||||
}
|
||||
|
||||
@@ -11,12 +11,14 @@ import (
|
||||
|
||||
// MobileDependency collect all dependencies for mobile platform
|
||||
type MobileDependency struct {
|
||||
// Android only
|
||||
TunAdapter device.TunAdapter
|
||||
IFaceDiscover stdnet.ExternalIFaceDiscover
|
||||
// Android and iOS
|
||||
NetworkChangeListener listener.NetworkChangeListener
|
||||
HostDNSAddresses []netip.AddrPort
|
||||
DnsReadyListener dns.ReadyListener
|
||||
|
||||
// Android only
|
||||
TunAdapter device.TunAdapter
|
||||
IFaceDiscover stdnet.ExternalIFaceDiscover
|
||||
HostDNSAddresses []netip.AddrPort
|
||||
DnsReadyListener dns.ReadyListener
|
||||
|
||||
// iOS only
|
||||
DnsManager dns.IosDnsManager
|
||||
|
||||
@@ -95,7 +95,7 @@ func (d *DnsInterceptor) RemoveRoute() error {
|
||||
|
||||
// AllowedIPs should use real IPs
|
||||
if d.currentPeerKey != "" {
|
||||
if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil {
|
||||
if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err))
|
||||
}
|
||||
}
|
||||
@@ -172,7 +172,7 @@ func (d *DnsInterceptor) removeAllowedIP(realPrefix netip.Prefix) error {
|
||||
}
|
||||
|
||||
// AllowedIPs use real IPs
|
||||
if _, err := d.allowedIPsRefcounter.Decrement(realPrefix); err != nil {
|
||||
if _, err := d.allowedIPsRefcounter.Decrement(realPrefix, d.currentPeerKey); err != nil {
|
||||
return fmt.Errorf("remove allowed IP %s: %v", realPrefix, err)
|
||||
}
|
||||
|
||||
@@ -205,7 +205,7 @@ func (d *DnsInterceptor) RemoveAllowedIPs() error {
|
||||
for _, prefixes := range d.interceptedDomains {
|
||||
for _, prefix := range prefixes {
|
||||
// AllowedIPs use real IPs
|
||||
if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil {
|
||||
if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -135,7 +135,7 @@ func (r *Route) RemoveAllowedIPs() error {
|
||||
var merr *multierror.Error
|
||||
for _, domainPrefixes := range r.dynamicDomains {
|
||||
for _, prefix := range domainPrefixes {
|
||||
if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil {
|
||||
if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err))
|
||||
}
|
||||
}
|
||||
@@ -320,7 +320,7 @@ func (r *Route) removeRoutes(prefixes []netip.Prefix) ([]netip.Prefix, error) {
|
||||
merr = multierror.Append(merr, fmt.Errorf("remove dynamic route for IP %s: %w", prefix, err))
|
||||
}
|
||||
if r.currentPeerKey != "" {
|
||||
if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil {
|
||||
if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,6 +52,10 @@ type Manager interface {
|
||||
UpdateRoutes(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error
|
||||
ClassifyRoutes(newRoutes []*route.Route) (map[route.ID]*route.Route, route.HAMap)
|
||||
TriggerSelection(route.HAMap)
|
||||
SelectRoutes(ids []route.NetID, appendRoute bool) error
|
||||
DeselectRoutes(ids []route.NetID) error
|
||||
SelectAllRoutes()
|
||||
DeselectAllRoutes()
|
||||
GetRouteSelector() *routeselector.RouteSelector
|
||||
GetClientRoutes() route.HAMap
|
||||
GetSelectedClientRoutes() route.HAMap
|
||||
@@ -216,7 +220,7 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) {
|
||||
)
|
||||
}
|
||||
|
||||
m.allowedIPsRefCounter = refcounter.New(
|
||||
m.allowedIPsRefCounter = refcounter.NewAllowedIPs(
|
||||
func(prefix netip.Prefix, peerKey string) (string, error) {
|
||||
// save peerKey to use it in the remove function
|
||||
return peerKey, m.wgInterface.AddAllowedIP(peerKey, prefix)
|
||||
@@ -800,7 +804,7 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI
|
||||
var info exitNodeInfo
|
||||
|
||||
for haID, routes := range clientRoutes {
|
||||
if !m.isExitNodeRoute(routes) {
|
||||
if !isExitNodeRoutes(routes) {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -820,13 +824,6 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI
|
||||
return info
|
||||
}
|
||||
|
||||
func (m *DefaultManager) isExitNodeRoute(routes []*route.Route) bool {
|
||||
if len(routes) == 0 {
|
||||
return false
|
||||
}
|
||||
return route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network)
|
||||
}
|
||||
|
||||
func (m *DefaultManager) categorizeUserSelection(netID route.NetID, info *exitNodeInfo) {
|
||||
if m.routeSelector.IsSelected(netID) {
|
||||
info.userSelected = append(info.userSelected, netID)
|
||||
|
||||
@@ -16,6 +16,8 @@ type MockManager struct {
|
||||
ClassifyRoutesFunc func(routes []*route.Route) (map[route.ID]*route.Route, route.HAMap)
|
||||
UpdateRoutesFunc func(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error
|
||||
TriggerSelectionFunc func(haMap route.HAMap)
|
||||
SelectRoutesFunc func(ids []route.NetID, appendRoute bool) error
|
||||
DeselectRoutesFunc func(ids []route.NetID) error
|
||||
GetRouteSelectorFunc func() *routeselector.RouteSelector
|
||||
GetClientRoutesFunc func() route.HAMap
|
||||
GetSelectedClientRoutesFunc func() route.HAMap
|
||||
@@ -55,6 +57,30 @@ func (m *MockManager) TriggerSelection(networks route.HAMap) {
|
||||
}
|
||||
}
|
||||
|
||||
// SelectRoutes mock implementation of SelectRoutes from Manager interface
|
||||
func (m *MockManager) SelectRoutes(ids []route.NetID, appendRoute bool) error {
|
||||
if m.SelectRoutesFunc != nil {
|
||||
return m.SelectRoutesFunc(ids, appendRoute)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeselectRoutes mock implementation of DeselectRoutes from Manager interface
|
||||
func (m *MockManager) DeselectRoutes(ids []route.NetID) error {
|
||||
if m.DeselectRoutesFunc != nil {
|
||||
return m.DeselectRoutesFunc(ids)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SelectAllRoutes mock implementation of SelectAllRoutes from Manager interface
|
||||
func (m *MockManager) SelectAllRoutes() {
|
||||
}
|
||||
|
||||
// DeselectAllRoutes mock implementation of DeselectAllRoutes from Manager interface
|
||||
func (m *MockManager) DeselectAllRoutes() {
|
||||
}
|
||||
|
||||
// GetRouteSelector mock implementation of GetRouteSelector from Manager interface
|
||||
func (m *MockManager) GetRouteSelector() *routeselector.RouteSelector {
|
||||
if m.GetRouteSelectorFunc != nil {
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
package notifier
|
||||
|
||||
import (
|
||||
"container/list"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sort"
|
||||
@@ -16,20 +15,12 @@ import (
|
||||
|
||||
type Notifier struct {
|
||||
mu sync.Mutex
|
||||
cond *sync.Cond
|
||||
currentPrefixes []string
|
||||
listener listener.NetworkChangeListener
|
||||
queue *list.List
|
||||
closed bool
|
||||
}
|
||||
|
||||
func NewNotifier() *Notifier {
|
||||
n := &Notifier{
|
||||
queue: list.New(),
|
||||
}
|
||||
n.cond = sync.NewCond(&n.mu)
|
||||
go n.deliverLoop()
|
||||
return n
|
||||
return &Notifier{}
|
||||
}
|
||||
|
||||
func (n *Notifier) SetListener(listener listener.NetworkChangeListener) {
|
||||
@@ -59,44 +50,19 @@ func (n *Notifier) OnNewPrefixes(prefixes []netip.Prefix) {
|
||||
sort.Strings(newNets)
|
||||
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
if slices.Equal(n.currentPrefixes, newNets) {
|
||||
n.mu.Unlock()
|
||||
return
|
||||
}
|
||||
n.currentPrefixes = newNets
|
||||
routes := strings.Join(n.currentPrefixes, ",")
|
||||
n.queue.PushBack(routes)
|
||||
n.cond.Signal()
|
||||
n.mu.Unlock()
|
||||
if n.listener != nil {
|
||||
n.listener.OnNetworkChanged(strings.Join(n.currentPrefixes, ","))
|
||||
}
|
||||
}
|
||||
|
||||
func (n *Notifier) Close() {
|
||||
n.mu.Lock()
|
||||
n.closed = true
|
||||
n.cond.Signal()
|
||||
n.mu.Unlock()
|
||||
}
|
||||
|
||||
func (n *Notifier) GetInitialRouteRanges() []string {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (n *Notifier) deliverLoop() {
|
||||
for {
|
||||
n.mu.Lock()
|
||||
for n.queue.Len() == 0 && !n.closed {
|
||||
n.cond.Wait()
|
||||
}
|
||||
if n.closed && n.queue.Len() == 0 {
|
||||
n.mu.Unlock()
|
||||
return
|
||||
}
|
||||
routes := n.queue.Remove(n.queue.Front()).(string)
|
||||
l := n.listener
|
||||
n.mu.Unlock()
|
||||
|
||||
if l != nil {
|
||||
l.OnNetworkChanged(routes)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -54,7 +54,7 @@ func (m *reconcileWGMock) GetNet() *netstack.Net { return n
|
||||
func TestReconcilePeerAllowedIPs(t *testing.T) {
|
||||
wg := &reconcileWGMock{}
|
||||
m := &DefaultManager{wgInterface: wg}
|
||||
m.allowedIPsRefCounter = refcounter.New[netip.Prefix, string, string](
|
||||
m.allowedIPsRefCounter = refcounter.NewAllowedIPs(
|
||||
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
|
||||
func(netip.Prefix, string) error { return nil },
|
||||
)
|
||||
|
||||
206
client/internal/routemanager/refcounter/allowedips.go
Normal file
206
client/internal/routemanager/refcounter/allowedips.go
Normal file
@@ -0,0 +1,206 @@
|
||||
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
|
||||
}
|
||||
241
client/internal/routemanager/refcounter/allowedips_test.go
Normal file
241
client/internal/routemanager/refcounter/allowedips_test.go
Normal file
@@ -0,0 +1,241 @@
|
||||
package refcounter
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeWG models WireGuard's cryptokey routing: a prefix can be installed on exactly one peer.
|
||||
// failAdd/failRemove make the next add/remove fail once, to exercise the self-healing error paths.
|
||||
type fakeWG struct {
|
||||
installed map[netip.Prefix]string
|
||||
adds int
|
||||
removes int
|
||||
failAdd bool
|
||||
failRemove bool
|
||||
}
|
||||
|
||||
func newFakeWG() *fakeWG {
|
||||
return &fakeWG{installed: map[netip.Prefix]string{}}
|
||||
}
|
||||
|
||||
func (f *fakeWG) counter() *AllowedIPsRefCounter {
|
||||
return NewAllowedIPs(
|
||||
func(prefix netip.Prefix, peerKey string) (string, error) {
|
||||
if f.failAdd {
|
||||
f.failAdd = false
|
||||
return "", errors.New("add failed")
|
||||
}
|
||||
f.adds++
|
||||
f.installed[prefix] = peerKey
|
||||
return peerKey, nil
|
||||
},
|
||||
func(prefix netip.Prefix, peerKey string) error {
|
||||
if f.failRemove {
|
||||
f.failRemove = false
|
||||
return errors.New("remove failed")
|
||||
}
|
||||
f.removes++
|
||||
// only clear if this peer is the one installed, mirroring wg semantics
|
||||
if f.installed[prefix] == peerKey {
|
||||
delete(f.installed, prefix)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func mustPrefix(t *testing.T, s string) netip.Prefix {
|
||||
t.Helper()
|
||||
p, err := netip.ParsePrefix(s)
|
||||
if err != nil {
|
||||
t.Fatalf("parse prefix %q: %v", s, err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func mustIncrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] {
|
||||
t.Helper()
|
||||
ref, err := c.Increment(p, peer)
|
||||
if err != nil {
|
||||
t.Fatalf("Increment(%v, %s): %v", p, peer, err)
|
||||
}
|
||||
return ref
|
||||
}
|
||||
|
||||
func mustDecrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] {
|
||||
t.Helper()
|
||||
ref, err := c.Decrement(p, peer)
|
||||
if err != nil {
|
||||
t.Fatalf("Decrement(%v, %s): %v", p, peer, err)
|
||||
}
|
||||
return ref
|
||||
}
|
||||
|
||||
// TestAllowedIPs_SwapOnActivePeerRemoval reproduces the reported bug: two networks with the same
|
||||
// prefix routed by different peers. Removing the network whose peer is installed must hand the
|
||||
// prefix over to the surviving peer instead of leaving it on the removed one.
|
||||
func TestAllowedIPs_SwapOnActivePeerRemoval(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p := mustPrefix(t, "10.44.8.0/24")
|
||||
|
||||
mustIncrement(t, c, p, "peerA")
|
||||
mustIncrement(t, c, p, "peerB")
|
||||
// First peer wins while both are present.
|
||||
if got := f.installed[p]; got != "peerA" {
|
||||
t.Fatalf("expected peerA installed, got %q", got)
|
||||
}
|
||||
|
||||
// Remove the active peer's network -> must swap to peerB.
|
||||
mustDecrement(t, c, p, "peerA")
|
||||
if got := f.installed[p]; got != "peerB" {
|
||||
t.Fatalf("BUG: prefix stuck on removed peer, want peerB got %q", got)
|
||||
}
|
||||
|
||||
// Remove the last one -> prefix gone.
|
||||
mustDecrement(t, c, p, "peerB")
|
||||
if _, ok := f.installed[p]; ok {
|
||||
t.Fatalf("expected prefix removed, still installed on %q", f.installed[p])
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowedIPs_RemoveNonActivePeer removing a non-installed peer must not touch WireGuard.
|
||||
func TestAllowedIPs_RemoveNonActivePeer(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p := mustPrefix(t, "10.44.8.0/24")
|
||||
|
||||
mustIncrement(t, c, p, "peerA")
|
||||
mustIncrement(t, c, p, "peerB")
|
||||
removesBefore := f.removes
|
||||
|
||||
mustDecrement(t, c, p, "peerB")
|
||||
if f.installed[p] != "peerA" {
|
||||
t.Fatalf("active peer must stay peerA, got %q", f.installed[p])
|
||||
}
|
||||
if f.removes != removesBefore {
|
||||
t.Fatalf("removing a non-active peer must not call wg remove")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowedIPs_SamePeerMultipleRefs two references via the same peer must keep the prefix until
|
||||
// the last reference is released (the reason the per-peer count must be an int, not a set).
|
||||
func TestAllowedIPs_SamePeerMultipleRefs(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p := mustPrefix(t, "10.44.8.0/24")
|
||||
|
||||
mustIncrement(t, c, p, "peerA")
|
||||
mustIncrement(t, c, p, "peerA")
|
||||
if f.adds != 1 {
|
||||
t.Fatalf("expected a single wg add for the same peer, got %d", f.adds)
|
||||
}
|
||||
|
||||
mustDecrement(t, c, p, "peerA")
|
||||
if f.installed[p] != "peerA" {
|
||||
t.Fatalf("prefix must stay while a reference remains, got %q", f.installed[p])
|
||||
}
|
||||
if f.removes != 0 {
|
||||
t.Fatalf("no wg remove expected while a reference remains, got %d", f.removes)
|
||||
}
|
||||
|
||||
mustDecrement(t, c, p, "peerA")
|
||||
if _, ok := f.installed[p]; ok {
|
||||
t.Fatalf("prefix must be removed after last reference")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowedIPs_RefCountAndActive checks the Ref returned to callers (used for the HA-disabled log).
|
||||
func TestAllowedIPs_RefCountAndActive(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p := mustPrefix(t, "10.44.8.0/24")
|
||||
|
||||
ref := mustIncrement(t, c, p, "peerA")
|
||||
if ref.Count != 1 || ref.Out != "peerA" {
|
||||
t.Fatalf("want {1, peerA}, got {%d, %q}", ref.Count, ref.Out)
|
||||
}
|
||||
ref = mustIncrement(t, c, p, "peerB")
|
||||
if ref.Count != 2 || ref.Out != "peerA" {
|
||||
t.Fatalf("want {2, peerA}, got {%d, %q}", ref.Count, ref.Out)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowedIPs_Flush removes everything installed and clears the counter.
|
||||
func TestAllowedIPs_Flush(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p1 := mustPrefix(t, "10.44.8.0/24")
|
||||
p2 := mustPrefix(t, "10.44.9.0/24")
|
||||
|
||||
mustIncrement(t, c, p1, "peerA")
|
||||
mustIncrement(t, c, p2, "peerB")
|
||||
|
||||
if err := c.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(f.installed) != 0 {
|
||||
t.Fatalf("expected all prefixes removed, got %v", f.installed)
|
||||
}
|
||||
// After flush, a fresh increment must add again.
|
||||
mustIncrement(t, c, p1, "peerC")
|
||||
if f.installed[p1] != "peerC" {
|
||||
t.Fatalf("counter not reset after flush")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowedIPs_SelfHealAfterSwapAddError ensures a failed add during a swap does not permanently
|
||||
// strand the prefix: the next Decrement (or Increment) must retry and install a surviving peer.
|
||||
func TestAllowedIPs_SelfHealAfterSwapAddError(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p := mustPrefix(t, "10.44.8.0/24")
|
||||
|
||||
mustIncrement(t, c, p, "peerA")
|
||||
mustIncrement(t, c, p, "peerB")
|
||||
mustIncrement(t, c, p, "peerC")
|
||||
|
||||
// Removing the active peerA triggers a swap to a survivor; make the add fail once.
|
||||
f.failAdd = true
|
||||
if _, err := c.Decrement(p, "peerA"); err == nil {
|
||||
t.Fatalf("expected error from failed swap add")
|
||||
}
|
||||
if _, ok := f.installed[p]; ok {
|
||||
t.Fatalf("nothing should be installed after a failed swap add, got %q", f.installed[p])
|
||||
}
|
||||
|
||||
// A later Decrement of a non-active survivor must retry the hand-off (self-heal), not stay stuck.
|
||||
ref := mustDecrement(t, c, p, "peerC")
|
||||
if got := f.installed[p]; got == "" {
|
||||
t.Fatalf("self-heal failed: prefix left unrouted after add recovered")
|
||||
}
|
||||
if ref.Out == "" {
|
||||
t.Fatalf("expected an active peer after self-heal, got empty")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowedIPs_SelfHealAfterRemoveError ensures a failed remove during a swap is retried instead
|
||||
// of leaving e.active stuck on a peer that no longer holds references.
|
||||
func TestAllowedIPs_SelfHealAfterRemoveError(t *testing.T) {
|
||||
f := newFakeWG()
|
||||
c := f.counter()
|
||||
p := mustPrefix(t, "10.44.8.0/24")
|
||||
|
||||
mustIncrement(t, c, p, "peerA")
|
||||
mustIncrement(t, c, p, "peerB")
|
||||
|
||||
// Releasing active peerA must detach it (remove) then add peerB; fail the remove once.
|
||||
f.failRemove = true
|
||||
if _, err := c.Decrement(p, "peerA"); err == nil {
|
||||
t.Fatalf("expected error from failed remove")
|
||||
}
|
||||
|
||||
// Next Decrement of the non-active survivor retries: removes stale peerA, installs peerB.
|
||||
mustDecrement(t, c, p, "peerB")
|
||||
// peerB had only one ref, so after retry the prefix is fully released.
|
||||
if _, ok := f.installed[p]; ok {
|
||||
t.Fatalf("expected prefix released after self-heal, still on %q", f.installed[p])
|
||||
}
|
||||
}
|
||||
@@ -5,5 +5,7 @@ import "net/netip"
|
||||
// RouteRefCounter is a Counter for Route, it doesn't take any input on Increment and doesn't use any output on Decrement
|
||||
type RouteRefCounter = Counter[netip.Prefix, struct{}, struct{}]
|
||||
|
||||
// AllowedIPsRefCounter is a Counter for AllowedIPs, it takes a peer key on Increment and passes it back to Decrement
|
||||
type AllowedIPsRefCounter = Counter[netip.Prefix, string, string]
|
||||
// AllowedIPsRefCounter tracks WireGuard AllowedIPs per prefix. Unlike the generic Counter it is peer-aware:
|
||||
// a prefix can be claimed by several peers at once and WireGuard allows a given prefix on exactly one peer,
|
||||
// so the counter records the per-peer reference count and swaps the installed peer when the active one is released.
|
||||
// See allowedips.go.
|
||||
|
||||
138
client/internal/routemanager/selection.go
Normal file
138
client/internal/routemanager/selection.go
Normal file
@@ -0,0 +1,138 @@
|
||||
package routemanager
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
"github.com/hashicorp/go-multierror"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
// SelectRoutes selects the routes with the given network IDs and applies the
|
||||
// new selection. V4/v6 exit-node pairs are expanded automatically. Exit nodes
|
||||
// are mutually exclusive: if the selection activates an exit node, every other
|
||||
// available exit node is deselected so two can't be active at once. With
|
||||
// appendRoute=false the previous selection is replaced instead of extended.
|
||||
func (m *DefaultManager) SelectRoutes(ids []route.NetID, appendRoute bool) error {
|
||||
if err := m.selectRoutes(ids, appendRoute); err != nil {
|
||||
return err
|
||||
}
|
||||
m.TriggerSelection(m.GetClientRoutes())
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeselectRoutes removes the routes with the given network IDs from the
|
||||
// selection and applies the change. V4/v6 exit-node pairs are expanded
|
||||
// automatically.
|
||||
func (m *DefaultManager) DeselectRoutes(ids []route.NetID) error {
|
||||
if err := m.deselectRoutes(ids); err != nil {
|
||||
return err
|
||||
}
|
||||
m.TriggerSelection(m.GetClientRoutes())
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *DefaultManager) deselectRoutes(ids []route.NetID) error {
|
||||
routesMap := m.GetClientRoutesWithNetID()
|
||||
routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap)
|
||||
|
||||
log.Debugf("deselecting routes with ids: %v", routes)
|
||||
|
||||
if err := m.routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil {
|
||||
return fmt.Errorf("deselect routes: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SelectAllRoutes selects every available route and applies the selection.
|
||||
// Exit nodes stay mutually exclusive: at most one remains active.
|
||||
func (m *DefaultManager) SelectAllRoutes() {
|
||||
m.selectAllRoutes()
|
||||
m.TriggerSelection(m.GetClientRoutes())
|
||||
}
|
||||
|
||||
func (m *DefaultManager) selectAllRoutes() {
|
||||
m.routeSelector.SelectAllRoutes()
|
||||
|
||||
// Select-all wipes every explicit selection, so exit nodes fall back to
|
||||
// management's auto-apply flags — which may mark several at once.
|
||||
// Reconcile immediately so at most one exit node stays active instead of
|
||||
// waiting for the next network map to enforce it.
|
||||
m.mux.Lock()
|
||||
defer m.mux.Unlock()
|
||||
m.updateRouteSelectorFromManagement(m.clientRoutes)
|
||||
}
|
||||
|
||||
// DeselectAllRoutes deselects every route and applies the change.
|
||||
func (m *DefaultManager) DeselectAllRoutes() {
|
||||
m.routeSelector.DeselectAllRoutes()
|
||||
m.TriggerSelection(m.GetClientRoutes())
|
||||
}
|
||||
|
||||
func (m *DefaultManager) selectRoutes(ids []route.NetID, appendRoute bool) error {
|
||||
routesMap := m.GetClientRoutesWithNetID()
|
||||
routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap)
|
||||
allIDs := maps.Keys(routesMap)
|
||||
|
||||
log.Debugf("selecting routes with ids: %v", routes)
|
||||
|
||||
// A partial failure (e.g. an unknown ID in the request) still selects the
|
||||
// valid routes, so exclusivity below must run regardless of the error.
|
||||
var merr *multierror.Error
|
||||
if err := m.routeSelector.SelectRoutes(routes, appendRoute, allIDs); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("select routes: %w", err))
|
||||
}
|
||||
|
||||
// Exit nodes are mutually exclusive: if this selection activates an
|
||||
// exit node, deselect every other available exit node so two can't be
|
||||
// selected at once. Non-exit route selections are left untouched.
|
||||
if requestActivatesExitNode(routes, routesMap) {
|
||||
if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 {
|
||||
if err := m.routeSelector.DeselectRoutes(others, allIDs); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("deselect sibling exit nodes: %w", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nberrors.FormatErrorOrNil(merr)
|
||||
}
|
||||
|
||||
func isExitNodeRoutes(routes []*route.Route) bool {
|
||||
return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network))
|
||||
}
|
||||
|
||||
// requestActivatesExitNode reports whether any requested NetID maps to an exit
|
||||
// node (default route) in the current route table.
|
||||
func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool {
|
||||
for _, id := range requested {
|
||||
if isExitNodeRoutes(routesMap[id]) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// otherExitNodeIDs returns every available exit-node NetID that is not in the
|
||||
// requested set — the siblings to deselect so a single exit node stays active.
|
||||
func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID {
|
||||
keep := make(map[route.NetID]struct{}, len(requested))
|
||||
for _, id := range requested {
|
||||
keep[id] = struct{}{}
|
||||
}
|
||||
var others []route.NetID
|
||||
for id, routes := range routesMap {
|
||||
if !isExitNodeRoutes(routes) {
|
||||
continue
|
||||
}
|
||||
if _, ok := keep[id]; ok {
|
||||
continue
|
||||
}
|
||||
others = append(others, id)
|
||||
}
|
||||
return others
|
||||
}
|
||||
129
client/internal/routemanager/selection_test.go
Normal file
129
client/internal/routemanager/selection_test.go
Normal file
@@ -0,0 +1,129 @@
|
||||
package routemanager
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/routeselector"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
func v6ExitRoute(netID, peer string) *route.Route {
|
||||
return &route.Route{
|
||||
NetID: route.NetID(netID),
|
||||
Network: netip.MustParsePrefix("::/0"),
|
||||
Peer: peer,
|
||||
}
|
||||
}
|
||||
|
||||
func newSelectionTestManager() *DefaultManager {
|
||||
return &DefaultManager{
|
||||
routeSelector: routeselector.NewRouteSelector(),
|
||||
clientRoutes: route.HAMap{
|
||||
"exitA|0.0.0.0/0": {exitRoute("exitA", "p1", true)},
|
||||
"exitA-v6|::/0": {v6ExitRoute("exitA-v6", "p1")},
|
||||
"exitB|0.0.0.0/0": {exitRoute("exitB", "p2", true)},
|
||||
"lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectRoutes_ExitNodeExclusivity(t *testing.T) {
|
||||
m := newSelectionTestManager()
|
||||
|
||||
// Selecting an exit node selects its v6 pair and deselects the sibling.
|
||||
require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true))
|
||||
assert.True(t, m.routeSelector.IsSelected("exitA"), "exitA should be selected")
|
||||
assert.True(t, m.routeSelector.IsSelected("exitA-v6"), "the v6 pair follows its v4 base")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitB"), "the sibling exit node must be deselected")
|
||||
|
||||
// Switching to the sibling deselects the previous exit node and its v6 pair.
|
||||
require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true))
|
||||
assert.True(t, m.routeSelector.IsSelected("exitB"), "exitB should now be selected")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitA"), "the previous exit node must be deselected")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "the previous exit node's v6 pair must be deselected")
|
||||
assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched")
|
||||
|
||||
// Selecting a non-exit route leaves the active exit node alone.
|
||||
require.NoError(t, m.selectRoutes([]route.NetID{"lan"}, true))
|
||||
assert.True(t, m.routeSelector.IsSelected("exitB"), "selecting a non-exit route keeps the exit node")
|
||||
|
||||
// Deselecting the active exit node turns every exit node off.
|
||||
require.NoError(t, m.deselectRoutes([]route.NetID{"exitB"}))
|
||||
assert.False(t, m.routeSelector.IsSelected("exitB"), "exitB should be deselected")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitA"), "exitA stays deselected")
|
||||
assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched")
|
||||
}
|
||||
|
||||
func TestSelectRoutes_PartialErrorStillEnforcesExclusivity(t *testing.T) {
|
||||
// The unknown ID must be reported, but the valid exit node in the same
|
||||
// request is still selected — so its sibling must still be deselected.
|
||||
// Both orderings are covered: processing must continue past the invalid
|
||||
// ID wherever it sits in the request.
|
||||
requests := map[string][]route.NetID{
|
||||
"invalid id first": {"missing", "exitB"},
|
||||
"invalid id last": {"exitB", "missing"},
|
||||
}
|
||||
|
||||
for name, ids := range requests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
m := newSelectionTestManager()
|
||||
|
||||
require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true))
|
||||
|
||||
err := m.selectRoutes(ids, true)
|
||||
assert.Error(t, err, "unknown id must be reported")
|
||||
assert.True(t, m.routeSelector.IsSelected("exitB"), "valid exit node from the request is selected")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitA"), "sibling exit node must be deselected despite the error")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "sibling's v6 pair must be deselected too")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectAllRoutes_KeepsSingleExitNode(t *testing.T) {
|
||||
// Both exit nodes are marked for auto-apply by management
|
||||
// (SkipAutoApply=false), the state where select-all could turn on two at
|
||||
// once without the immediate reconciliation.
|
||||
m := &DefaultManager{
|
||||
routeSelector: routeselector.NewRouteSelector(),
|
||||
clientRoutes: route.HAMap{
|
||||
"exitA|0.0.0.0/0": {exitRoute("exitA", "p1", false)},
|
||||
"exitB|0.0.0.0/0": {exitRoute("exitB", "p2", false)},
|
||||
"lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}},
|
||||
},
|
||||
}
|
||||
|
||||
require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true))
|
||||
|
||||
m.selectAllRoutes()
|
||||
|
||||
assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit routes are all selected")
|
||||
assert.True(t, m.routeSelector.IsSelected("exitA"), "the deterministic management pick stays active")
|
||||
assert.False(t, m.routeSelector.IsSelected("exitB"), "select-all must not leave a second exit node active")
|
||||
}
|
||||
|
||||
func TestSelectRoutes_UnknownRoute(t *testing.T) {
|
||||
m := newSelectionTestManager()
|
||||
|
||||
assert.Error(t, m.selectRoutes([]route.NetID{"missing"}, true), "selecting an unavailable route must fail")
|
||||
assert.Error(t, m.deselectRoutes([]route.NetID{"missing"}), "deselecting an unavailable route must fail")
|
||||
}
|
||||
|
||||
func TestExitNodeSelectionHelpers(t *testing.T) {
|
||||
routesMap := map[route.NetID][]*route.Route{
|
||||
"exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}},
|
||||
"exitB": {{Network: netip.MustParsePrefix("::/0")}},
|
||||
"lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}},
|
||||
}
|
||||
|
||||
assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node")
|
||||
assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node")
|
||||
assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node")
|
||||
assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node")
|
||||
|
||||
others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"})
|
||||
assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored")
|
||||
}
|
||||
@@ -15,6 +15,11 @@ type Route struct {
|
||||
route *route.Route
|
||||
routeRefCounter *refcounter.RouteRefCounter
|
||||
allowedIPsRefcounter *refcounter.AllowedIPsRefCounter
|
||||
// currentPeerKey is the routing peer this watcher currently has the prefix installed on
|
||||
// (the HA winner elected by the watcher). It can differ from route.Peer and change on
|
||||
// failover, so it is recorded on AddAllowedIPs and used on RemoveAllowedIPs to decrement
|
||||
// the exact peer that was incremented.
|
||||
currentPeerKey string
|
||||
}
|
||||
|
||||
func NewRoute(params common.HandlerParams) *Route {
|
||||
@@ -52,12 +57,15 @@ func (r *Route) AddAllowedIPs(peerKey string) error {
|
||||
ref.Out,
|
||||
)
|
||||
}
|
||||
r.currentPeerKey = peerKey
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Route) RemoveAllowedIPs() error {
|
||||
if _, err := r.allowedIPsRefcounter.Decrement(r.route.Network); err != nil {
|
||||
return err
|
||||
var err error
|
||||
if _, decErr := r.allowedIPsRefcounter.Decrement(r.route.Network, r.currentPeerKey); decErr != nil {
|
||||
err = fmt.Errorf("remove allowed IP %s: %w", r.route.Network, decErr)
|
||||
}
|
||||
return nil
|
||||
r.currentPeerKey = ""
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -20,6 +20,8 @@ const (
|
||||
rpFilterPath = "net.ipv4.conf.all.rp_filter"
|
||||
rpFilterInterfacePath = "net.ipv4.conf.%s.rp_filter"
|
||||
srcValidMarkPath = "net.ipv4.conf.all.src_valid_mark"
|
||||
percentEscape = "%25"
|
||||
dotEscape = "%2E"
|
||||
)
|
||||
|
||||
type iface interface {
|
||||
@@ -56,7 +58,11 @@ func Setup(wgIface iface) (map[string]int, error) {
|
||||
continue
|
||||
}
|
||||
|
||||
i := fmt.Sprintf(rpFilterInterfacePath, intf.Name)
|
||||
// Escape '%' and '.' so they survive the dot-to-slash conversion in Set()
|
||||
safeName := strings.ReplaceAll(intf.Name, "%", percentEscape)
|
||||
safeName = strings.ReplaceAll(safeName, ".", dotEscape)
|
||||
|
||||
i := fmt.Sprintf(rpFilterInterfacePath, safeName)
|
||||
oldVal, err := Set(i, 2, true)
|
||||
if err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
@@ -70,7 +76,11 @@ func Setup(wgIface iface) (map[string]int, error) {
|
||||
|
||||
// Set sets a sysctl configuration, if onlyIfOne is true it will only set the new value if it's set to 1
|
||||
func Set(key string, desiredValue int, onlyIfOne bool) (int, error) {
|
||||
path := fmt.Sprintf("/proc/sys/%s", strings.ReplaceAll(key, ".", "/"))
|
||||
path := strings.ReplaceAll(key, ".", "/")
|
||||
// Unescape interface dots and percent signs
|
||||
path = strings.ReplaceAll(path, dotEscape, ".")
|
||||
path = strings.ReplaceAll(path, percentEscape, "%")
|
||||
path = fmt.Sprintf("/proc/sys/%s", path)
|
||||
currentValue, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("read sysctl %s: %w", key, err)
|
||||
|
||||
124
client/internal/tunnelnotifier/notifier.go
Normal file
124
client/internal/tunnelnotifier/notifier.go
Normal file
@@ -0,0 +1,124 @@
|
||||
package tunnelnotifier
|
||||
|
||||
import (
|
||||
"container/list"
|
||||
"sync"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/dns"
|
||||
"github.com/netbirdio/netbird/client/internal/listener"
|
||||
)
|
||||
|
||||
type eventKind int
|
||||
|
||||
const (
|
||||
eventRoutes eventKind = iota
|
||||
eventIfaceIP
|
||||
eventIfaceIPv6
|
||||
eventDNS
|
||||
)
|
||||
|
||||
var (
|
||||
_ listener.NetworkChangeListener = (*Notifier)(nil)
|
||||
_ dns.IosDnsManager = (*Notifier)(nil)
|
||||
)
|
||||
|
||||
type event struct {
|
||||
kind eventKind
|
||||
payload string
|
||||
}
|
||||
|
||||
type Notifier struct {
|
||||
mu sync.Mutex
|
||||
cond *sync.Cond
|
||||
queue *list.List
|
||||
closed bool
|
||||
done chan struct{}
|
||||
|
||||
listener listener.NetworkChangeListener
|
||||
dnsManager dns.IosDnsManager
|
||||
}
|
||||
|
||||
func New(l listener.NetworkChangeListener, dm dns.IosDnsManager) *Notifier {
|
||||
n := &Notifier{
|
||||
queue: list.New(),
|
||||
done: make(chan struct{}),
|
||||
listener: l,
|
||||
dnsManager: dm,
|
||||
}
|
||||
n.cond = sync.NewCond(&n.mu)
|
||||
go n.deliverLoop()
|
||||
return n
|
||||
}
|
||||
|
||||
func (n *Notifier) OnNetworkChanged(routes string) {
|
||||
n.enqueue(event{kind: eventRoutes, payload: routes})
|
||||
}
|
||||
|
||||
func (n *Notifier) SetInterfaceIP(ip string) {
|
||||
n.enqueue(event{kind: eventIfaceIP, payload: ip})
|
||||
}
|
||||
|
||||
func (n *Notifier) SetInterfaceIPv6(ip string) {
|
||||
n.enqueue(event{kind: eventIfaceIPv6, payload: ip})
|
||||
}
|
||||
|
||||
func (n *Notifier) ApplyDns(config string) {
|
||||
n.enqueue(event{kind: eventDNS, payload: config})
|
||||
}
|
||||
|
||||
// Close stops accepting new events and blocks until the delivery loop has
|
||||
// drained all queued events and exited.
|
||||
func (n *Notifier) Close() {
|
||||
n.mu.Lock()
|
||||
n.closed = true
|
||||
n.cond.Signal()
|
||||
n.mu.Unlock()
|
||||
<-n.done
|
||||
}
|
||||
|
||||
func (n *Notifier) enqueue(ev event) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
if n.closed {
|
||||
return
|
||||
}
|
||||
n.queue.PushBack(ev)
|
||||
n.cond.Signal()
|
||||
}
|
||||
|
||||
func (n *Notifier) deliverLoop() {
|
||||
defer close(n.done)
|
||||
for {
|
||||
n.mu.Lock()
|
||||
for n.queue.Len() == 0 && !n.closed {
|
||||
n.cond.Wait()
|
||||
}
|
||||
if n.closed && n.queue.Len() == 0 {
|
||||
n.mu.Unlock()
|
||||
return
|
||||
}
|
||||
ev := n.queue.Remove(n.queue.Front()).(event)
|
||||
l := n.listener
|
||||
dm := n.dnsManager
|
||||
n.mu.Unlock()
|
||||
|
||||
switch ev.kind {
|
||||
case eventRoutes:
|
||||
if l != nil {
|
||||
l.OnNetworkChanged(ev.payload)
|
||||
}
|
||||
case eventIfaceIP:
|
||||
if l != nil {
|
||||
l.SetInterfaceIP(ev.payload)
|
||||
}
|
||||
case eventIfaceIPv6:
|
||||
if l != nil {
|
||||
l.SetInterfaceIPv6(ev.payload)
|
||||
}
|
||||
case eventDNS:
|
||||
if dm != nil {
|
||||
dm.ApplyDns(ev.payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
192
client/internal/tunnelnotifier/notifier_test.go
Normal file
192
client/internal/tunnelnotifier/notifier_test.go
Normal file
@@ -0,0 +1,192 @@
|
||||
package tunnelnotifier
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type call struct {
|
||||
kind string
|
||||
payload string
|
||||
}
|
||||
|
||||
type recorder struct {
|
||||
mu sync.Mutex
|
||||
calls []call
|
||||
inFlight atomic.Int32
|
||||
overlap atomic.Bool
|
||||
delay time.Duration
|
||||
}
|
||||
|
||||
func (r *recorder) record(kind, payload string) {
|
||||
if r.inFlight.Add(1) != 1 {
|
||||
r.overlap.Store(true)
|
||||
}
|
||||
if r.delay > 0 {
|
||||
time.Sleep(r.delay)
|
||||
}
|
||||
r.mu.Lock()
|
||||
r.calls = append(r.calls, call{kind: kind, payload: payload})
|
||||
r.mu.Unlock()
|
||||
r.inFlight.Add(-1)
|
||||
}
|
||||
|
||||
func (r *recorder) count() int {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return len(r.calls)
|
||||
}
|
||||
|
||||
func (r *recorder) snapshot() []call {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
out := make([]call, len(r.calls))
|
||||
copy(out, r.calls)
|
||||
return out
|
||||
}
|
||||
|
||||
type fakeListener struct {
|
||||
rec *recorder
|
||||
}
|
||||
|
||||
func (f *fakeListener) OnNetworkChanged(routes string) {
|
||||
f.rec.record("routes", routes)
|
||||
}
|
||||
|
||||
func (f *fakeListener) SetInterfaceIP(ip string) {
|
||||
f.rec.record("ip", ip)
|
||||
}
|
||||
|
||||
func (f *fakeListener) SetInterfaceIPv6(ip string) {
|
||||
f.rec.record("ipv6", ip)
|
||||
}
|
||||
|
||||
type fakeDNSManager struct {
|
||||
rec *recorder
|
||||
}
|
||||
|
||||
func (f *fakeDNSManager) ApplyDns(config string) {
|
||||
f.rec.record("dns", config)
|
||||
}
|
||||
|
||||
func TestFIFOOrder(t *testing.T) {
|
||||
rec := &recorder{}
|
||||
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
|
||||
defer n.Close()
|
||||
|
||||
n.SetInterfaceIP("10.0.0.1")
|
||||
n.SetInterfaceIPv6("fd00::1")
|
||||
n.ApplyDns(`{"domains":[]}`)
|
||||
n.OnNetworkChanged("10.0.0.0/8,192.168.0.0/16")
|
||||
n.ApplyDns(`{"domains":["example.com"]}`)
|
||||
|
||||
require.Eventually(t, func() bool { return rec.count() == 5 }, time.Second, time.Millisecond)
|
||||
|
||||
expected := []call{
|
||||
{kind: "ip", payload: "10.0.0.1"},
|
||||
{kind: "ipv6", payload: "fd00::1"},
|
||||
{kind: "dns", payload: `{"domains":[]}`},
|
||||
{kind: "routes", payload: "10.0.0.0/8,192.168.0.0/16"},
|
||||
{kind: "dns", payload: `{"domains":["example.com"]}`},
|
||||
}
|
||||
assert.Equal(t, expected, rec.snapshot())
|
||||
}
|
||||
|
||||
func TestNoOverlappingCalls(t *testing.T) {
|
||||
rec := &recorder{delay: 100 * time.Microsecond}
|
||||
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
|
||||
defer n.Close()
|
||||
|
||||
const producers = 8
|
||||
const perProducer = 25
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < producers; i++ {
|
||||
wg.Add(1)
|
||||
go func(id int) {
|
||||
defer wg.Done()
|
||||
for j := 0; j < perProducer; j++ {
|
||||
payload := fmt.Sprintf("%d-%d", id, j)
|
||||
switch j % 4 {
|
||||
case 0:
|
||||
n.OnNetworkChanged(payload)
|
||||
case 1:
|
||||
n.SetInterfaceIP(payload)
|
||||
case 2:
|
||||
n.SetInterfaceIPv6(payload)
|
||||
case 3:
|
||||
n.ApplyDns(payload)
|
||||
}
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
require.Eventually(t, func() bool { return rec.count() == producers*perProducer }, 5*time.Second, time.Millisecond)
|
||||
assert.False(t, rec.overlap.Load())
|
||||
}
|
||||
|
||||
func TestDNSAndRoutesInterleaved(t *testing.T) {
|
||||
rec := &recorder{delay: 100 * time.Microsecond}
|
||||
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
|
||||
defer n.Close()
|
||||
|
||||
const events = 50
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < events; i++ {
|
||||
n.ApplyDns(fmt.Sprintf("dns-%d", i))
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < events; i++ {
|
||||
n.OnNetworkChanged(fmt.Sprintf("routes-%d", i))
|
||||
}
|
||||
}()
|
||||
wg.Wait()
|
||||
|
||||
require.Eventually(t, func() bool { return rec.count() == 2*events }, 5*time.Second, time.Millisecond)
|
||||
assert.False(t, rec.overlap.Load())
|
||||
|
||||
var dnsSeen, routesSeen int
|
||||
for _, c := range rec.snapshot() {
|
||||
switch c.kind {
|
||||
case "dns":
|
||||
assert.Equal(t, fmt.Sprintf("dns-%d", dnsSeen), c.payload)
|
||||
dnsSeen++
|
||||
case "routes":
|
||||
assert.Equal(t, fmt.Sprintf("routes-%d", routesSeen), c.payload)
|
||||
routesSeen++
|
||||
}
|
||||
}
|
||||
assert.Equal(t, events, dnsSeen)
|
||||
assert.Equal(t, events, routesSeen)
|
||||
}
|
||||
|
||||
func TestCloseDrainsQueue(t *testing.T) {
|
||||
rec := &recorder{delay: time.Millisecond}
|
||||
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
|
||||
|
||||
const events = 20
|
||||
for i := 0; i < events; i++ {
|
||||
n.OnNetworkChanged(fmt.Sprintf("routes-%d", i))
|
||||
}
|
||||
n.Close()
|
||||
|
||||
require.Equal(t, events, rec.count(), "Close must not return before all queued events are delivered")
|
||||
|
||||
n.OnNetworkChanged("after-close")
|
||||
n.ApplyDns("after-close")
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
assert.Equal(t, events, rec.count())
|
||||
}
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/internal/auth"
|
||||
@@ -637,23 +636,18 @@ func (c *Client) SelectRoute(id string) error {
|
||||
}
|
||||
|
||||
routeManager := engine.GetRouteManager()
|
||||
routeSelector := routeManager.GetRouteSelector()
|
||||
if id == "All" {
|
||||
log.Debugf("select all routes")
|
||||
routeSelector.SelectAllRoutes()
|
||||
} else {
|
||||
log.Debugf("select route with id: %s", id)
|
||||
routes := toNetIDs([]string{id})
|
||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||
routes = route.ExpandV6ExitPairs(routes, routesMap)
|
||||
if err := routeSelector.SelectRoutes(routes, true, maps.Keys(routesMap)); err != nil {
|
||||
log.Debugf("error when selecting routes: %s", err)
|
||||
return fmt.Errorf("select routes: %w", err)
|
||||
}
|
||||
routeManager.SelectAllRoutes()
|
||||
return nil
|
||||
}
|
||||
routeManager.TriggerSelection(routeManager.GetClientRoutes())
|
||||
return nil
|
||||
|
||||
log.Debugf("select route with id: %s", id)
|
||||
if err := routeManager.SelectRoutes(toNetIDs([]string{id}), true); err != nil {
|
||||
log.Debugf("error when selecting routes: %s", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) DeselectRoute(id string) error {
|
||||
@@ -667,21 +661,17 @@ func (c *Client) DeselectRoute(id string) error {
|
||||
}
|
||||
|
||||
routeManager := engine.GetRouteManager()
|
||||
routeSelector := routeManager.GetRouteSelector()
|
||||
if id == "All" {
|
||||
log.Debugf("deselect all routes")
|
||||
routeSelector.DeselectAllRoutes()
|
||||
} else {
|
||||
log.Debugf("deselect route with id: %s", id)
|
||||
routes := toNetIDs([]string{id})
|
||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||
routes = route.ExpandV6ExitPairs(routes, routesMap)
|
||||
if err := routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil {
|
||||
log.Debugf("error when deselecting routes: %s", err)
|
||||
return fmt.Errorf("deselect routes: %w", err)
|
||||
}
|
||||
routeManager.DeselectAllRoutes()
|
||||
return nil
|
||||
}
|
||||
|
||||
log.Debugf("deselect route with id: %s", id)
|
||||
if err := routeManager.DeselectRoutes(toNetIDs([]string{id})); err != nil {
|
||||
log.Debugf("error when deselecting routes: %s", err)
|
||||
return err
|
||||
}
|
||||
routeManager.TriggerSelection(routeManager.GetClientRoutes())
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/exp/maps"
|
||||
"google.golang.org/grpc/codes"
|
||||
gstatus "google.golang.org/grpc/status"
|
||||
|
||||
@@ -161,30 +160,11 @@ func (s *Server) SelectNetworks(_ context.Context, req *proto.SelectNetworksRequ
|
||||
return nil, fmt.Errorf("no route manager")
|
||||
}
|
||||
|
||||
routeSelector := routeManager.GetRouteSelector()
|
||||
if req.GetAll() {
|
||||
routeSelector.SelectAllRoutes()
|
||||
} else {
|
||||
routes := toNetIDs(req.GetNetworkIDs())
|
||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||
routes = route.ExpandV6ExitPairs(routes, routesMap)
|
||||
netIdRoutes := maps.Keys(routesMap)
|
||||
if err := routeSelector.SelectRoutes(routes, req.GetAppend(), netIdRoutes); err != nil {
|
||||
return nil, fmt.Errorf("select routes: %w", err)
|
||||
}
|
||||
|
||||
// Exit nodes are mutually exclusive: if this selection activates an
|
||||
// exit node, deselect every other available exit node so two can't be
|
||||
// selected at once. Non-exit route selections are left untouched.
|
||||
if requestActivatesExitNode(routes, routesMap) {
|
||||
if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 {
|
||||
if err := routeSelector.DeselectRoutes(others, netIdRoutes); err != nil {
|
||||
return nil, fmt.Errorf("deselect sibling exit nodes: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
routeManager.SelectAllRoutes()
|
||||
} else if err := routeManager.SelectRoutes(toNetIDs(req.GetNetworkIDs()), req.GetAppend()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
routeManager.TriggerSelection(routeManager.GetClientRoutes())
|
||||
|
||||
s.statusRecorder.PublishEvent(
|
||||
proto.SystemEvent_INFO,
|
||||
@@ -224,19 +204,11 @@ func (s *Server) DeselectNetworks(_ context.Context, req *proto.SelectNetworksRe
|
||||
return nil, fmt.Errorf("no route manager")
|
||||
}
|
||||
|
||||
routeSelector := routeManager.GetRouteSelector()
|
||||
if req.GetAll() {
|
||||
routeSelector.DeselectAllRoutes()
|
||||
} else {
|
||||
routes := toNetIDs(req.GetNetworkIDs())
|
||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||
routes = route.ExpandV6ExitPairs(routes, routesMap)
|
||||
netIdRoutes := maps.Keys(routesMap)
|
||||
if err := routeSelector.DeselectRoutes(routes, netIdRoutes); err != nil {
|
||||
return nil, fmt.Errorf("deselect routes: %w", err)
|
||||
}
|
||||
routeManager.DeselectAllRoutes()
|
||||
} else if err := routeManager.DeselectRoutes(toNetIDs(req.GetNetworkIDs())); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
routeManager.TriggerSelection(routeManager.GetClientRoutes())
|
||||
|
||||
s.statusRecorder.PublishEvent(
|
||||
proto.SystemEvent_INFO,
|
||||
@@ -261,37 +233,3 @@ func toNetIDs(routes []string) []route.NetID {
|
||||
return netIDs
|
||||
}
|
||||
|
||||
func isExitNodeRoutes(routes []*route.Route) bool {
|
||||
return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network))
|
||||
}
|
||||
|
||||
// requestActivatesExitNode reports whether any requested NetID maps to an exit
|
||||
// node (default route) in the current route table.
|
||||
func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool {
|
||||
for _, id := range requested {
|
||||
if isExitNodeRoutes(routesMap[id]) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// otherExitNodeIDs returns every available exit-node NetID that is not in the
|
||||
// requested set — the siblings to deselect so a single exit node stays active.
|
||||
func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID {
|
||||
keep := make(map[route.NetID]struct{}, len(requested))
|
||||
for _, id := range requested {
|
||||
keep[id] = struct{}{}
|
||||
}
|
||||
var others []route.NetID
|
||||
for id, routes := range routesMap {
|
||||
if !isExitNodeRoutes(routes) {
|
||||
continue
|
||||
}
|
||||
if _, ok := keep[id]; ok {
|
||||
continue
|
||||
}
|
||||
others = append(others, id)
|
||||
}
|
||||
return others
|
||||
}
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
func TestExitNodeSelectionHelpers(t *testing.T) {
|
||||
routesMap := map[route.NetID][]*route.Route{
|
||||
"exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}},
|
||||
"exitB": {{Network: netip.MustParsePrefix("::/0")}},
|
||||
"lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}},
|
||||
}
|
||||
|
||||
assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node")
|
||||
assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node")
|
||||
assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node")
|
||||
assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node")
|
||||
|
||||
others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"})
|
||||
assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored")
|
||||
}
|
||||
@@ -2,6 +2,7 @@ import { type ReactNode } from "react";
|
||||
import { useTranslation } from "react-i18next";
|
||||
import { Browser } from "@wailsio/runtime";
|
||||
import { DownloadIcon, NotepadText } from "lucide-react";
|
||||
import { Update as UpdateSvc } from "@bindings/services";
|
||||
import { Button } from "@/components/buttons/Button";
|
||||
import { useClientVersion } from "@/contexts/ClientVersionContext";
|
||||
import { cn } from "@/lib/cn";
|
||||
@@ -14,6 +15,12 @@ function openUrl(url: string) {
|
||||
});
|
||||
}
|
||||
|
||||
function openInstallerDownload() {
|
||||
UpdateSvc.DownloadURL()
|
||||
.then(openUrl)
|
||||
.catch(() => openUrl(GITHUB_RELEASES));
|
||||
}
|
||||
|
||||
export function UpdateVersionCard() {
|
||||
const { t } = useTranslation();
|
||||
const { updateVersion, enforced, triggerUpdate } = useClientVersion();
|
||||
@@ -37,11 +44,7 @@ export function UpdateVersionCard() {
|
||||
{t("update.card.installNow")}
|
||||
</Button>
|
||||
) : (
|
||||
<Button
|
||||
variant={"primary"}
|
||||
size={"xs"}
|
||||
onClick={() => openUrl(GITHUB_RELEASES)}
|
||||
>
|
||||
<Button variant={"primary"} size={"xs"} onClick={openInstallerDownload}>
|
||||
<DownloadIcon size={14} />
|
||||
{t("update.card.getInstaller")}
|
||||
</Button>
|
||||
|
||||
@@ -281,6 +281,9 @@ func newApplication(onSecondInstance func()) *application.App {
|
||||
Linux: application.LinuxOptions{
|
||||
ProgramName: "netbird",
|
||||
},
|
||||
Windows: application.WindowsOptions{
|
||||
WndProcInterceptor: endSessionInterceptor(),
|
||||
},
|
||||
SingleInstance: &application.SingleInstanceOptions{
|
||||
UniqueID: "io.netbird.ui",
|
||||
OnSecondInstanceLaunch: func(_ application.SecondInstanceData) {
|
||||
@@ -367,6 +370,9 @@ func newMainWindow(app *application.App, prefStore *preferences.Store) *applicat
|
||||
|
||||
// Hide instead of quit on close; "really quit" is reached via tray -> Quit.
|
||||
window.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) {
|
||||
if services.ShuttingDown() {
|
||||
return
|
||||
}
|
||||
e.Cancel()
|
||||
window.Hide()
|
||||
})
|
||||
|
||||
24
client/ui/services/shutdown.go
Normal file
24
client/ui/services/shutdown.go
Normal file
@@ -0,0 +1,24 @@
|
||||
package services
|
||||
|
||||
import "sync/atomic"
|
||||
|
||||
var (
|
||||
sessionEnding atomic.Bool
|
||||
quitting atomic.Bool
|
||||
)
|
||||
|
||||
func BeginSessionEnd() {
|
||||
sessionEnding.Store(true)
|
||||
}
|
||||
|
||||
func AbortSessionEnd() {
|
||||
sessionEnding.Store(false)
|
||||
}
|
||||
|
||||
func BeginShutdown() {
|
||||
quitting.Store(true)
|
||||
}
|
||||
|
||||
func ShuttingDown() bool {
|
||||
return sessionEnding.Load() || quitting.Load()
|
||||
}
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
"github.com/netbirdio/netbird/client/ui/updater"
|
||||
"github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
// UpdateResult mirrors TriggerUpdateResponse.
|
||||
@@ -33,6 +34,12 @@ func (s *Update) GetState() updater.State {
|
||||
return s.holder.Get()
|
||||
}
|
||||
|
||||
// DownloadURL returns the platform-appropriate installer download link for
|
||||
// manual (non-enforced) updates.
|
||||
func (s *Update) DownloadURL() string {
|
||||
return version.DownloadUrl()
|
||||
}
|
||||
|
||||
// Quit exits the app. Scheduled off the calling goroutine so the JS caller's
|
||||
// response returns before the runtime tears down.
|
||||
func (s *Update) Quit() {
|
||||
|
||||
@@ -154,6 +154,9 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo
|
||||
})
|
||||
// Hide (not destroy) on close to keep React state; reset to General for a flash-free reopen.
|
||||
s.settings.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) {
|
||||
if ShuttingDown() {
|
||||
return
|
||||
}
|
||||
e.Cancel()
|
||||
s.app.Event.Emit(EventSettingsOpen, "general")
|
||||
s.settings.Hide()
|
||||
@@ -393,6 +396,9 @@ func (s *WindowManager) CloseWelcome() {
|
||||
// OpenError shows the custom error dialog; title/message are pre-localised and ride in the
|
||||
// start URL. A second error replaces the open one via SetURL. Singleton, destroyed on close.
|
||||
func (s *WindowManager) OpenError(title, message string) {
|
||||
if ShuttingDown() {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
startURL := errorDialogURL(title, message)
|
||||
|
||||
7
client/ui/shutdown_other.go
Normal file
7
client/ui/shutdown_other.go
Normal file
@@ -0,0 +1,7 @@
|
||||
//go:build !windows && !android && !ios && !freebsd && !js
|
||||
|
||||
package main
|
||||
|
||||
func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) {
|
||||
return nil
|
||||
}
|
||||
36
client/ui/shutdown_windows.go
Normal file
36
client/ui/shutdown_windows.go
Normal file
@@ -0,0 +1,36 @@
|
||||
//go:build windows
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/ui/services"
|
||||
)
|
||||
|
||||
const (
|
||||
wmQueryEndSession = 0x0011
|
||||
wmEndSession = 0x0016
|
||||
)
|
||||
|
||||
func endSessionInterceptor() func(hwnd uintptr, msg uint32, wParam, lParam uintptr) (uintptr, bool) {
|
||||
return func(_ uintptr, msg uint32, wParam, _ uintptr) (uintptr, bool) {
|
||||
switch msg {
|
||||
case wmQueryEndSession:
|
||||
services.BeginSessionEnd()
|
||||
return 1, true
|
||||
case wmEndSession:
|
||||
if wParam == 0 {
|
||||
services.AbortSessionEnd()
|
||||
return 0, true
|
||||
}
|
||||
log.Info("windows session is ending; exiting immediately")
|
||||
os.Exit(0)
|
||||
return 0, true
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -32,9 +32,8 @@ const (
|
||||
|
||||
quitDownTimeout = 5 * time.Second
|
||||
|
||||
urlGitHubRepo = "https://github.com/netbirdio/netbird"
|
||||
urlGitHubReleases = "https://github.com/netbirdio/netbird/releases/latest"
|
||||
urlDocs = "https://docs.netbird.io"
|
||||
urlGitHubRepo = "https://github.com/netbirdio/netbird"
|
||||
urlDocs = "https://docs.netbird.io"
|
||||
)
|
||||
|
||||
// TrayServices bundles the services the tray menu needs, grouped so NewTray
|
||||
@@ -453,6 +452,7 @@ func (t *Tray) buildMenu() *application.Menu {
|
||||
}
|
||||
|
||||
func (t *Tray) handleQuit() {
|
||||
services.BeginShutdown()
|
||||
t.profileMu.Lock()
|
||||
if t.switchCancel != nil {
|
||||
t.switchCancel()
|
||||
|
||||
@@ -25,6 +25,9 @@ type sendFn func(notifications.NotificationOptions) error
|
||||
// event-dispatch goroutine that panic is fatal process-wide; recover() turns
|
||||
// it into a logged no-op.
|
||||
func safeSendNotification(send sendFn, what string, opts notifications.NotificationOptions) (err error) {
|
||||
if services.ShuttingDown() {
|
||||
return nil
|
||||
}
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
log.Errorf("notify %s: recovered from panic (notification bus unavailable): %v", what, r)
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/client/ui/services"
|
||||
"github.com/netbirdio/netbird/client/ui/updater"
|
||||
"github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
// trayUpdater owns the tray UI that reacts to auto-update. Composed inside Tray.
|
||||
@@ -76,15 +77,15 @@ func (u *trayUpdater) applyLanguage() {
|
||||
u.refreshMenuItem(state)
|
||||
}
|
||||
|
||||
// handleClick opens the GitHub releases page when not Enforced, otherwise shows
|
||||
// the progress page and asks the daemon to start the installer.
|
||||
// handleClick opens the installer download link when not Enforced, otherwise
|
||||
// shows the progress page and asks the daemon to start the installer.
|
||||
func (u *trayUpdater) handleClick() {
|
||||
u.mu.Lock()
|
||||
state := u.state
|
||||
u.mu.Unlock()
|
||||
|
||||
if !state.Enforced {
|
||||
_ = u.app.Browser.OpenURL(urlGitHubReleases)
|
||||
_ = u.app.Browser.OpenURL(version.DownloadUrl())
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -109,7 +109,7 @@ sequenceDiagram
|
||||
Chk->>Inj: continue
|
||||
Inj->>Inj: inject NetBird identity headers per provider config
|
||||
Inj->>Grd: continue
|
||||
Grd->>Grd: enforce model allowlist
|
||||
Grd->>Grd: enforce per-provider allowlist (fail-closed backstop)
|
||||
Grd->>Up: forward (over WireGuard)
|
||||
Up-->>Resp: response (JSON or SSE stream)
|
||||
Resp->>Resp: parse usage tokens, completion
|
||||
@@ -135,6 +135,21 @@ sequenceDiagram
|
||||
(`redact_pii = settings.RedactPii`). Phones, emails, credit cards,
|
||||
PII names — see `redact.go` for the full set. See
|
||||
[`modules/31-proxy-middleware-builtin.md`](modules/31-proxy-middleware-builtin.md).
|
||||
- The model allowlist is enforced in TWO places. `CheckLLMPolicyLimits`
|
||||
is authoritative: it resolves the policy that governs this
|
||||
(provider, caller-groups) and denies (`deny_code = llm_policy.model_blocked`)
|
||||
when no applicable policy permits the model — so an allowlist scoped to
|
||||
one group/provider never leaks to another, and an un-guardrailed policy
|
||||
is genuinely unrestricted. `llm_guardrail` is a per-provider fail-closed
|
||||
backstop: it only carries an allowlist for a provider every authorising
|
||||
policy restricts, and blocks unknown/undetermined models even when
|
||||
management is unreachable. Because that backstop allowlist is the UNION
|
||||
of every restricting policy's models, per-group narrowing lives only in
|
||||
the authoritative check: during a `CheckLLMPolicyLimits` outage
|
||||
`llm_limit_check` fails open, so a caller can reach any model in the
|
||||
provider's union — a group scoped to model A could reach model B if
|
||||
another group restricts the same provider to B. This is the documented
|
||||
fail-open trade-off; a future flag may switch it to fail-closed.
|
||||
- SSE streaming requires special handling on the response side; the
|
||||
parser must handle partial chunks without buffering the whole
|
||||
stream. See [`modules/32-proxy-llm-parsers.md`](modules/32-proxy-llm-parsers.md).
|
||||
|
||||
@@ -122,7 +122,7 @@ At request time the path is independent: the proxy calls `SelectPolicyForRequest
|
||||
| on_request | 1 | `llm_router` | `{"providers":[{id, models[], upstream_*, auth_header_*, allowed_group_ids[]}]}` | **true** |
|
||||
| on_request | 2 | `llm_limit_check` | `{}` | – |
|
||||
| on_request | 3 | `llm_identity_inject` | `{"providers":[{provider_id, header_pair?, json_metadata?, extra_headers?}]}` | **true** |
|
||||
| on_request | 4 | `llm_guardrail` | `{"model_allowlist"?, "prompt_capture":{enabled,redact_pii}}` | – |
|
||||
| on_request | 4 | `llm_guardrail` | `{"provider_allowlists"?: {providerID: []model}, "prompt_capture":{enabled,redact_pii}}` | – |
|
||||
| on_response | 5 | `llm_limit_record` | `{}` (runs LAST at runtime) | – |
|
||||
| on_response | 6 | `cost_meter` | `{}` | – |
|
||||
| on_response | 7 | `llm_response_parser` | `{"capture_completion": <bool>, "redact_pii"?: true}` | – |
|
||||
|
||||
@@ -244,7 +244,7 @@ no mocks. Tests: `TestChain_AllowPath_StampsAttributionAndRecordsCounter`
|
||||
| `llm_router` | `{providers: [{id, models, upstream_scheme, upstream_host, upstream_path?, auth_header_name, auth_header_value, allowed_group_ids}]}` |
|
||||
| `llm_limit_check` | `{}` — pulls `MgmtClient` from `FactoryContext` |
|
||||
| `llm_identity_inject` | `{providers: [{provider_id, header_pair?|json_metadata?, extra_headers?}]}` |
|
||||
| `llm_guardrail` | `{model_allowlist: []string, prompt_capture: {enabled, redact_pii}}` |
|
||||
| `llm_guardrail` | `{provider_allowlists: {providerID: []string}, prompt_capture: {enabled, redact_pii}}` — allowlist keyed by resolved provider id; a provider absent from the map is unrestricted (fail-closed backstop; authoritative per-policy/group check is management's `CheckLLMPolicyLimits`) |
|
||||
| `llm_response_parser` | `{redact_pii?, capture_completion?: *bool}` |
|
||||
| `cost_meter` | `{pricing_path?}` (basename inside data-dir; defaults `pricing.yaml`) |
|
||||
| `llm_limit_record` | `{}` — same pattern as `llm_limit_check` |
|
||||
|
||||
209
e2e/agentnetwork/guardrail_block_test.go
Normal file
209
e2e/agentnetwork/guardrail_block_test.go
Normal file
@@ -0,0 +1,209 @@
|
||||
//go:build e2e
|
||||
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/e2e/harness"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// pathRoutedGuardrailCase is one provider's self-contained scenario: its own
|
||||
// provider, its own guardrail whose allowlist holds ONLY that provider's
|
||||
// allowed model, and its own policy. Each case runs in isolation (its own
|
||||
// proxy + client), so the guardrail the proxy enforces contains exactly this
|
||||
// provider's model — never a mixed cross-provider list.
|
||||
type pathRoutedGuardrailCase struct {
|
||||
name string
|
||||
catalogID string // agent-network catalog provider id
|
||||
wire string // harness.WireVertex | harness.WireBedrock
|
||||
allowEntry string // the single model id put on the guardrail allowlist
|
||||
allowModel string // model id sent that MUST be served (200)
|
||||
blockModel string // model id sent that MUST be denied (403 model_blocked)
|
||||
}
|
||||
|
||||
// TestGuardrailBlocksUnselectedModel_PathRouted is the end-to-end regression
|
||||
// guard for the customer report that a model-allowlist guardrail attached to a
|
||||
// policy has no effect for PATH-ROUTED providers — where the model travels in
|
||||
// the URL, not the JSON body: Google Vertex (…/models/{model}:rawPredict) and
|
||||
// AWS Bedrock (/model/{id}/invoke).
|
||||
//
|
||||
// Each provider is tested in isolation with a guardrail allowlisting a single
|
||||
// model of its own: the allowed model (in the URL path) is served (200) and an
|
||||
// unselected model (in the URL path) is denied 403 by the guardrail
|
||||
// (llm_policy.model_blocked) before the upstream. The Vertex case mirrors the
|
||||
// customer verbatim — allow Sonnet, and the unselected model is the exact
|
||||
// claude-opus-4-6 they reported reaching the model unblocked. The Bedrock case
|
||||
// sends a region-prefixed, versioned inference-profile id so URL-path model
|
||||
// normalization is exercised too.
|
||||
//
|
||||
// The provider is catch-all (no models), so the router forwards any model and a
|
||||
// 403 can only come from the guardrail, never model_not_routable. Only the
|
||||
// upstream LLM is mocked (the vLLM nginx answers any path with 200); management
|
||||
// synth/reconcile, the proxy middleware chain (URL-path model extraction,
|
||||
// router, guardrail) and the tunnel are all real, and the guardrail denies
|
||||
// before the upstream is dialed so the mock cannot influence the block. A
|
||||
// static bearer api key is used so the router injects a static Authorization
|
||||
// header instead of minting a GCP token — the only reason path-routed providers
|
||||
// normally need live credentials — so the test runs with none and is always on.
|
||||
func TestGuardrailBlocksUnselectedModel_PathRouted(t *testing.T) {
|
||||
cases := []pathRoutedGuardrailCase{
|
||||
{
|
||||
name: "vertex",
|
||||
catalogID: "vertex_ai_api",
|
||||
wire: harness.WireVertex,
|
||||
allowEntry: "claude-sonnet-4-5",
|
||||
allowModel: "claude-sonnet-4-5",
|
||||
blockModel: "claude-opus-4-6", // the customer-reported model
|
||||
},
|
||||
{
|
||||
name: "bedrock",
|
||||
catalogID: "bedrock_api",
|
||||
wire: harness.WireBedrock,
|
||||
allowEntry: "anthropic.claude-sonnet-4-5", // normalized catalog id
|
||||
allowModel: "us.anthropic.claude-sonnet-4-5-v1:0",
|
||||
blockModel: "us.anthropic.claude-opus-4-8-v1:0",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
runPathRoutedGuardrailCase(t, tc)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func runPathRoutedGuardrailCase(t *testing.T, tc pathRoutedGuardrailCase) {
|
||||
t.Helper()
|
||||
|
||||
const (
|
||||
vertexProject = "e2e-project"
|
||||
vertexRegion = "global"
|
||||
)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
vllm, err := harness.StartVLLM(ctx, srv)
|
||||
require.NoError(t, err, "start mock upstream")
|
||||
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
|
||||
|
||||
grp, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-" + tc.name})
|
||||
require.NoError(t, err, "create group")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grp.Id) })
|
||||
|
||||
ephemeral := false
|
||||
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
|
||||
Name: "e2e-guardrail-" + tc.name + "-client",
|
||||
Type: "reusable",
|
||||
ExpiresIn: 86400,
|
||||
UsageLimit: 0,
|
||||
AutoGroups: []string{grp.Id},
|
||||
Ephemeral: &ephemeral,
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
require.NotEmpty(t, sk.Key, "setup key plaintext")
|
||||
|
||||
// Catch-all provider (no models) so the router forwards any model; a static
|
||||
// bearer key means the router injects a static auth header instead of minting
|
||||
// a GCP token. Bootstraps the cluster if it isn't already.
|
||||
staticKey := "static-e2e-token"
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: tc.name,
|
||||
ProviderId: tc.catalogID,
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
})
|
||||
require.NoError(t, err, "create %s provider", tc.name)
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
// Guardrail allowlisting ONLY this provider's allowed model.
|
||||
var gr api.AgentNetworkGuardrailRequest
|
||||
gr.Name = "e2e-guardrail-" + tc.name
|
||||
gr.Checks.ModelAllowlist.Enabled = true
|
||||
gr.Checks.ModelAllowlist.Models = []string{tc.allowEntry}
|
||||
guard, err := srv.CreateGuardrail(ctx, gr)
|
||||
require.NoError(t, err, "create guardrail")
|
||||
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), guard.Id) })
|
||||
|
||||
enabled := true
|
||||
pol, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-guardrail-" + tc.name,
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grp.Id},
|
||||
DestinationProviderIds: []string{prov.Id},
|
||||
GuardrailIds: &[]string{guard.Id},
|
||||
})
|
||||
require.NoError(t, err, "create policy")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), pol.Id) })
|
||||
|
||||
settings, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "read settings")
|
||||
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
|
||||
|
||||
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-"+tc.name+"-proxy")
|
||||
require.NoError(t, err, "mint proxy token")
|
||||
px, err := harness.StartProxy(ctx, srv, proxyToken)
|
||||
require.NoError(t, err, "start proxy")
|
||||
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
|
||||
|
||||
cl, err := harness.StartClient(ctx, srv, sk.Key)
|
||||
require.NoError(t, err, "start client")
|
||||
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
|
||||
|
||||
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
|
||||
// Probe first: the GET resolves the endpoint and its first packet wakes the
|
||||
// lazy proxy peer, so WaitProxyPeer then observes it connected.
|
||||
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
|
||||
require.NoError(t, err, "resolve endpoint to proxy IP")
|
||||
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
|
||||
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
|
||||
}
|
||||
|
||||
send := func(model string) (int, string) {
|
||||
var code int
|
||||
var body string
|
||||
var cerr error
|
||||
switch tc.wire {
|
||||
case harness.WireVertex:
|
||||
code, body, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, vertexProject, vertexRegion, model, "Reply with exactly: pong", "")
|
||||
case harness.WireBedrock:
|
||||
code, body, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, model, "Reply with exactly: pong", "")
|
||||
default:
|
||||
t.Fatalf("unsupported wire %q", tc.wire)
|
||||
}
|
||||
require.NoError(t, cerr, "request must reach the proxy for %s", tc.name)
|
||||
return code, body
|
||||
}
|
||||
|
||||
// Allowed model (in the URL path) is served. Retry to absorb tunnel/DNS
|
||||
// jitter on the first call over the freshly warmed tunnel.
|
||||
var code int
|
||||
var body string
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
code, body = send(tc.allowModel)
|
||||
if code == 200 {
|
||||
break
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
assert.Equal(t, 200, code,
|
||||
"allowed %s model (URL path) must be served; body: %s\n=== proxy logs ===\n%s", tc.name, body, px.Logs(context.Background()))
|
||||
|
||||
// Unselected model (in the URL path) must be blocked by the guardrail.
|
||||
code, body = send(tc.blockModel)
|
||||
assert.Equal(t, 403, code,
|
||||
"unselected %s model (URL path) must be denied, not served; body: %s\n=== proxy logs ===\n%s", tc.name, body, px.Logs(context.Background()))
|
||||
assert.Contains(t, body, "llm_policy.model_blocked",
|
||||
"%s denial must come from the guardrail allowlist, not routing; body: %s", tc.name, body)
|
||||
}
|
||||
205
e2e/agentnetwork/guardrail_groupswitch_test.go
Normal file
205
e2e/agentnetwork/guardrail_groupswitch_test.go
Normal file
@@ -0,0 +1,205 @@
|
||||
//go:build e2e
|
||||
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/e2e/harness"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// TestGuardrailGroupSwitchTakesEffectAfterTTL proves that moving a peer between
|
||||
// groups flips its model-allowlist decision once the proxy's tunnel-peer cache
|
||||
// expires. The peer's groups reach the guardrail via ValidateTunnelPeer, which
|
||||
// the proxy caches; the switch is invisible until that cache expires. The proxy
|
||||
// runs with a short NB_PROXY_TUNNEL_CACHE_TTL so the flip happens in seconds
|
||||
// instead of the 5-minute default.
|
||||
//
|
||||
// Setup: one catch-all provider declaring modelA + modelB; polA (grpA -> allow
|
||||
// modelA) and polB (grpB -> allow modelB). The client starts in grpA. modelA is
|
||||
// served and modelB denied; after switching the client grpA -> grpB, modelB is
|
||||
// served and modelA denied. The cross-group deny comes from management's
|
||||
// per-policy/group CheckLLMPolicyLimits (the proxy backstop carries the union).
|
||||
func TestGuardrailGroupSwitchTakesEffectAfterTTL(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
const (
|
||||
modelA = "e2e-model-a"
|
||||
modelB = "e2e-model-b"
|
||||
// Short tunnel-cache TTL so a group switch propagates in seconds.
|
||||
// Exercises the NB_PROXY_TUNNEL_CACHE_TTL override.
|
||||
cacheTTL = 3 * time.Second
|
||||
)
|
||||
|
||||
vllm, err := harness.StartVLLM(ctx, srv)
|
||||
require.NoError(t, err, "start mock upstream")
|
||||
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
|
||||
|
||||
grpA, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-gswitch-a"})
|
||||
require.NoError(t, err, "create group A")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpA.Id) })
|
||||
|
||||
grpB, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-gswitch-b"})
|
||||
require.NoError(t, err, "create group B")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpB.Id) })
|
||||
|
||||
ephemeral := false
|
||||
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
|
||||
Name: "e2e-gswitch-client",
|
||||
Type: "reusable",
|
||||
ExpiresIn: 86400,
|
||||
UsageLimit: 0,
|
||||
AutoGroups: []string{grpA.Id}, // client starts in group A
|
||||
Ephemeral: &ephemeral,
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
require.NotEmpty(t, sk.Key, "setup key plaintext")
|
||||
|
||||
staticKey := "static-e2e-token"
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "gswitch",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: &[]api.AgentNetworkProviderModel{
|
||||
{Id: modelA, InputPer1k: 0.001, OutputPer1k: 0.001},
|
||||
{Id: modelB, InputPer1k: 0.001, OutputPer1k: 0.001},
|
||||
},
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
})
|
||||
require.NoError(t, err, "create provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
mkGuard := func(name, model string) api.AgentNetworkGuardrail {
|
||||
var gr api.AgentNetworkGuardrailRequest
|
||||
gr.Name = name
|
||||
gr.Checks.ModelAllowlist.Enabled = true
|
||||
gr.Checks.ModelAllowlist.Models = []string{model}
|
||||
g, gerr := srv.CreateGuardrail(ctx, gr)
|
||||
require.NoError(t, gerr, "create guardrail %s", name)
|
||||
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) })
|
||||
return g
|
||||
}
|
||||
gA := mkGuard("e2e-gswitch-a", modelA)
|
||||
gB := mkGuard("e2e-gswitch-b", modelB)
|
||||
|
||||
enabled := true
|
||||
polA, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-gswitch-a",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpA.Id},
|
||||
DestinationProviderIds: []string{prov.Id},
|
||||
GuardrailIds: &[]string{gA.Id},
|
||||
})
|
||||
require.NoError(t, err, "create policy A")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polA.Id) })
|
||||
|
||||
polB, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-gswitch-b",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpB.Id},
|
||||
DestinationProviderIds: []string{prov.Id},
|
||||
GuardrailIds: &[]string{gB.Id},
|
||||
})
|
||||
require.NoError(t, err, "create policy B")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polB.Id) })
|
||||
|
||||
settings, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "read settings")
|
||||
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
|
||||
|
||||
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-gswitch-proxy")
|
||||
require.NoError(t, err, "mint proxy token")
|
||||
px, err := harness.StartProxy(ctx, srv, proxyToken, map[string]string{
|
||||
"NB_PROXY_TUNNEL_CACHE_TTL": cacheTTL.String(),
|
||||
})
|
||||
require.NoError(t, err, "start proxy")
|
||||
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
|
||||
|
||||
cl, err := harness.StartClient(ctx, srv, sk.Key)
|
||||
require.NoError(t, err, "start client")
|
||||
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
|
||||
|
||||
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
|
||||
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
|
||||
require.NoError(t, err, "resolve endpoint to proxy IP")
|
||||
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
|
||||
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
|
||||
}
|
||||
|
||||
send := func(model string) (int, string) {
|
||||
code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "")
|
||||
require.NoError(t, cerr, "request must reach the proxy")
|
||||
return code, body
|
||||
}
|
||||
sendUntil := func(model string, want int, timeout time.Duration) (int, string) {
|
||||
var code int
|
||||
var body string
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
code, body = send(model)
|
||||
if code == want {
|
||||
return code, body
|
||||
}
|
||||
time.Sleep(2 * time.Second)
|
||||
}
|
||||
return code, body
|
||||
}
|
||||
|
||||
// Phase 1 — client is in group A: modelA served, modelB denied.
|
||||
code, body := sendUntil(modelA, 200, 90*time.Second)
|
||||
assert.Equal(t, 200, code,
|
||||
"group-A model must be served while the client is in group A; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
code, body = send(modelB)
|
||||
assert.Equal(t, 403, code,
|
||||
"group-B model must be denied while the client is in group A; body: %s", body)
|
||||
assert.Contains(t, body, "llm_policy.model_blocked")
|
||||
|
||||
// Switch the client peer from group A to group B.
|
||||
peerID := clientPeerInGroup(t, ctx, grpA.Id)
|
||||
_, err = srv.API().Groups.Update(ctx, grpB.Id, api.PutApiGroupsGroupIdJSONRequestBody{
|
||||
Name: grpB.Name,
|
||||
Peers: &[]string{peerID},
|
||||
})
|
||||
require.NoError(t, err, "add peer to group B")
|
||||
_, err = srv.API().Groups.Update(ctx, grpA.Id, api.PutApiGroupsGroupIdJSONRequestBody{
|
||||
Name: grpA.Name,
|
||||
Peers: &[]string{},
|
||||
})
|
||||
require.NoError(t, err, "remove peer from group A")
|
||||
|
||||
// Phase 2 — after the short TTL expires the proxy re-validates the peer,
|
||||
// sees group B, and the decision flips. Poll to absorb TTL + re-validation.
|
||||
code, body = sendUntil(modelB, 200, 60*time.Second)
|
||||
assert.Equal(t, 200, code,
|
||||
"after the group switch + TTL, the group-B model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
code, body = send(modelA)
|
||||
assert.Equal(t, 403, code,
|
||||
"after the switch, the old group-A model must be denied; body: %s", body)
|
||||
assert.Contains(t, body, "llm_policy.model_blocked")
|
||||
}
|
||||
|
||||
// clientPeerInGroup returns the id of the single peer that is a member of the
|
||||
// given group — the test client. The proxy peer is never added to test groups.
|
||||
func clientPeerInGroup(t *testing.T, ctx context.Context, groupID string) string {
|
||||
t.Helper()
|
||||
peers, err := srv.API().Peers.List(ctx)
|
||||
require.NoError(t, err, "list peers")
|
||||
for _, p := range peers {
|
||||
for _, g := range p.Groups {
|
||||
if g.Id == groupID {
|
||||
return p.Id
|
||||
}
|
||||
}
|
||||
}
|
||||
t.Fatalf("no peer found in group %s", groupID)
|
||||
return ""
|
||||
}
|
||||
201
e2e/agentnetwork/guardrail_multipolicy_test.go
Normal file
201
e2e/agentnetwork/guardrail_multipolicy_test.go
Normal file
@@ -0,0 +1,201 @@
|
||||
//go:build e2e
|
||||
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/e2e/harness"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// TestGuardrailMultiPolicyModelAllowlist: modelSelected served (200), grpOther's
|
||||
// modelOther denied for the grpMain client (403 model_blocked, no cross-group
|
||||
// leak), and openModel on the un-guardrailed policy's provider served (200).
|
||||
func TestGuardrailMultiPolicyModelAllowlist(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
const (
|
||||
modelSelected = "e2e-selected"
|
||||
modelOther = "e2e-other"
|
||||
openModel = "e2e-open"
|
||||
)
|
||||
|
||||
vllm, err := harness.StartVLLM(ctx, srv)
|
||||
require.NoError(t, err, "start mock upstream")
|
||||
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
|
||||
|
||||
grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-main"})
|
||||
require.NoError(t, err, "create main group")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) })
|
||||
|
||||
grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-guardrail-mp-other"})
|
||||
require.NoError(t, err, "create other group")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) })
|
||||
|
||||
ephemeral := false
|
||||
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
|
||||
Name: "e2e-guardrail-mp-client",
|
||||
Type: "reusable",
|
||||
ExpiresIn: 86400,
|
||||
UsageLimit: 0,
|
||||
AutoGroups: []string{grpMain.Id}, // client joins grpMain only
|
||||
Ephemeral: &ephemeral,
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
require.NotEmpty(t, sk.Key, "setup key plaintext")
|
||||
|
||||
staticKey := "static-e2e-token"
|
||||
models := func(ids ...string) *[]api.AgentNetworkProviderModel {
|
||||
out := make([]api.AgentNetworkProviderModel, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001})
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
// pRestricted declares the two guardrailed models so routing is deterministic
|
||||
// (model -> provider). Created first, so it carries the bootstrap cluster.
|
||||
pRestricted, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "restricted",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: models(modelSelected, modelOther),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
})
|
||||
require.NoError(t, err, "create restricted provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pRestricted.Id) })
|
||||
|
||||
pOpen, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "open",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: models(openModel),
|
||||
})
|
||||
require.NoError(t, err, "create open provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pOpen.Id) })
|
||||
|
||||
mkGuardrail := func(name, model string) api.AgentNetworkGuardrail {
|
||||
var gr api.AgentNetworkGuardrailRequest
|
||||
gr.Name = name
|
||||
gr.Checks.ModelAllowlist.Enabled = true
|
||||
gr.Checks.ModelAllowlist.Models = []string{model}
|
||||
g, gerr := srv.CreateGuardrail(ctx, gr)
|
||||
require.NoError(t, gerr, "create guardrail %s", name)
|
||||
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) })
|
||||
return g
|
||||
}
|
||||
gMain := mkGuardrail("e2e-guardrail-mp-main", modelSelected)
|
||||
gOther := mkGuardrail("e2e-guardrail-mp-other", modelOther)
|
||||
|
||||
enabled := true
|
||||
// polMain: grpMain restricted to modelSelected on pRestricted.
|
||||
polMain, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-guardrail-mp-main",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpMain.Id},
|
||||
DestinationProviderIds: []string{pRestricted.Id},
|
||||
GuardrailIds: &[]string{gMain.Id},
|
||||
})
|
||||
require.NoError(t, err, "create main policy")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) })
|
||||
|
||||
// polOther: grpOther restricted to modelOther on the SAME provider. The
|
||||
// client is not in grpOther, so modelOther must never be usable by it.
|
||||
polOther, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-guardrail-mp-other",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpOther.Id},
|
||||
DestinationProviderIds: []string{pRestricted.Id},
|
||||
GuardrailIds: &[]string{gOther.Id},
|
||||
})
|
||||
require.NoError(t, err, "create other policy")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) })
|
||||
|
||||
// polOpen: grpMain on pOpen with NO guardrail — unrestricted.
|
||||
polOpen, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-guardrail-mp-open",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpMain.Id},
|
||||
DestinationProviderIds: []string{pOpen.Id},
|
||||
})
|
||||
require.NoError(t, err, "create open policy")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOpen.Id) })
|
||||
|
||||
settings, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "read settings")
|
||||
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
|
||||
|
||||
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-guardrail-mp-proxy")
|
||||
require.NoError(t, err, "mint proxy token")
|
||||
px, err := harness.StartProxy(ctx, srv, proxyToken)
|
||||
require.NoError(t, err, "start proxy")
|
||||
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
|
||||
|
||||
cl, err := harness.StartClient(ctx, srv, sk.Key)
|
||||
require.NoError(t, err, "start client")
|
||||
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
|
||||
|
||||
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
|
||||
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
|
||||
require.NoError(t, err, "resolve endpoint to proxy IP")
|
||||
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
|
||||
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
|
||||
}
|
||||
|
||||
send := func(model string) (int, string) {
|
||||
code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "")
|
||||
require.NoError(t, cerr, "request must reach the proxy")
|
||||
return code, body
|
||||
}
|
||||
// sendUntil200 absorbs first-call tunnel/DNS jitter on the freshly warmed tunnel.
|
||||
sendUntil200 := func(model string) (int, string) {
|
||||
var code int
|
||||
var body string
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
code, body = send(model)
|
||||
if code == 200 {
|
||||
break
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
return code, body
|
||||
}
|
||||
|
||||
t.Run("selected model allowed for its group", func(t *testing.T) {
|
||||
code, body := sendUntil200(modelSelected)
|
||||
assert.Equal(t, 200, code,
|
||||
"grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
})
|
||||
|
||||
t.Run("other group's model does not leak", func(t *testing.T) {
|
||||
// modelOther is allowlisted only for grpOther. The grpMain client must be
|
||||
// denied by management's per-policy/group check — not waved through by an
|
||||
// account-wide union. This is the security-critical wrong-ALLOW guard.
|
||||
code, body := send(modelOther)
|
||||
assert.Equal(t, 403, code,
|
||||
"another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
assert.Contains(t, body, "llm_policy.model_blocked",
|
||||
"denial must be a model-allowlist decision; body: %s", body)
|
||||
})
|
||||
|
||||
t.Run("unguarded policy leaves its provider unrestricted", func(t *testing.T) {
|
||||
// polOpen carries no guardrail, so pOpen is unrestricted for grpMain. The
|
||||
// old account-wide union would have blocked openModel (it is on no
|
||||
// allowlist); it must now be served — the false-DENY guard.
|
||||
code, body := sendUntil200(openModel)
|
||||
assert.Equal(t, 200, code,
|
||||
"an un-guardrailed policy's provider must not be blocked by another policy's allowlist; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
})
|
||||
}
|
||||
422
e2e/agentnetwork/guardrail_pergroup_providers_test.go
Normal file
422
e2e/agentnetwork/guardrail_pergroup_providers_test.go
Normal file
@@ -0,0 +1,422 @@
|
||||
//go:build e2e
|
||||
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/e2e/harness"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// pergroupCase describes one provider surface for the per-group allowlist matrix.
|
||||
// selectedReq/otherReq are the model identifiers as they travel in the request
|
||||
// (URL path for Bedrock/Vertex, body "model" for chat/messages). selectedAllow/
|
||||
// otherAllow are the (normalized) forms the guardrail allowlist holds — for
|
||||
// Bedrock these differ from the request form so path normalization is exercised.
|
||||
type pergroupCase struct {
|
||||
name string
|
||||
catalogID string
|
||||
wire string // "chat", "messages", "vertex", "bedrock"
|
||||
models *[]api.AgentNetworkProviderModel
|
||||
selectedReq string
|
||||
selectedAllow string
|
||||
otherReq string
|
||||
otherAllow string
|
||||
|
||||
providerID string // filled during setup
|
||||
}
|
||||
|
||||
// TestGuardrailPerGroupAllowlist_AllProviders proves the per-policy/group model
|
||||
// allowlist end to end across every always-on provider surface, including the
|
||||
// path-routed ones (Vertex, Bedrock) where the model travels in the URL.
|
||||
//
|
||||
// For each provider two policies target it: grpMain (the client) is allowed only
|
||||
// selectedReq; grpOther (which the client is NOT in) is allowed only otherReq.
|
||||
// The client must get selectedReq served (200) and otherReq denied (403,
|
||||
// llm_policy.model_blocked) — the cross-group no-leak property. The deny is the
|
||||
// authoritative per-policy/group decision from management (the proxy per-provider
|
||||
// backstop carries the union of both models), so this also confirms management
|
||||
// receives the correct normalized model for path-routed providers.
|
||||
func TestGuardrailPerGroupAllowlist_AllProviders(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
const (
|
||||
vertexProject = "e2e-project"
|
||||
vertexRegion = "global"
|
||||
)
|
||||
|
||||
priced := func(ids ...string) *[]api.AgentNetworkProviderModel {
|
||||
out := make([]api.AgentNetworkProviderModel, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001})
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
cases := []*pergroupCase{
|
||||
{
|
||||
name: "openai", catalogID: "openai_api", wire: harness.WireChat,
|
||||
models: priced("oai-model-a", "oai-model-b"),
|
||||
selectedReq: "oai-model-a", selectedAllow: "oai-model-a",
|
||||
otherReq: "oai-model-b", otherAllow: "oai-model-b",
|
||||
},
|
||||
{
|
||||
name: "anthropic", catalogID: "anthropic_api", wire: harness.WireMessages,
|
||||
models: priced("ant-model-a", "ant-model-b"),
|
||||
selectedReq: "ant-model-a", selectedAllow: "ant-model-a",
|
||||
otherReq: "ant-model-b", otherAllow: "ant-model-b",
|
||||
},
|
||||
{
|
||||
// Vertex catalog ids travel bare in the rawPredict path.
|
||||
name: "vertex", catalogID: "vertex_ai_api", wire: "vertex",
|
||||
selectedReq: "claude-sonnet-4-5", selectedAllow: "claude-sonnet-4-5",
|
||||
otherReq: "claude-opus-4-6", otherAllow: "claude-opus-4-6",
|
||||
},
|
||||
{
|
||||
// Bedrock request ids are region-prefixed/versioned; the parser
|
||||
// normalizes them to the catalog key the allowlist holds.
|
||||
name: "bedrock", catalogID: "bedrock_api", wire: "bedrock",
|
||||
selectedReq: "us.anthropic.claude-sonnet-4-5-v1:0", selectedAllow: "anthropic.claude-sonnet-4-5",
|
||||
otherReq: "us.anthropic.claude-opus-4-8-v1:0", otherAllow: "anthropic.claude-opus-4-8",
|
||||
},
|
||||
}
|
||||
|
||||
vllm, err := harness.StartVLLM(ctx, srv)
|
||||
require.NoError(t, err, "start mock upstream")
|
||||
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
|
||||
|
||||
grpMain, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-pergroup-main"})
|
||||
require.NoError(t, err, "create main group")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpMain.Id) })
|
||||
|
||||
grpOther, err := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: "e2e-pergroup-other"})
|
||||
require.NoError(t, err, "create other group")
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), grpOther.Id) })
|
||||
|
||||
ephemeral := false
|
||||
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
|
||||
Name: "e2e-pergroup-client",
|
||||
Type: "reusable",
|
||||
ExpiresIn: 86400,
|
||||
UsageLimit: 0,
|
||||
AutoGroups: []string{grpMain.Id},
|
||||
Ephemeral: &ephemeral,
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
require.NotEmpty(t, sk.Key, "setup key plaintext")
|
||||
|
||||
staticKey := "static-e2e-token"
|
||||
enabled := true
|
||||
|
||||
for i, c := range cases {
|
||||
req := api.AgentNetworkProviderRequest{
|
||||
Name: "e2e-pergroup-" + c.name,
|
||||
ProviderId: c.catalogID,
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: c.models,
|
||||
}
|
||||
if i == 0 {
|
||||
req.BootstrapCluster = ptr(harness.AgentNetworkCluster)
|
||||
}
|
||||
prov, perr := srv.CreateProvider(ctx, req)
|
||||
require.NoError(t, perr, "create provider %s", c.name)
|
||||
c.providerID = prov.Id
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
gSel := mkAllowGuardrail(t, ctx, "e2e-pergroup-"+c.name+"-sel", c.selectedAllow)
|
||||
gOth := mkAllowGuardrail(t, ctx, "e2e-pergroup-"+c.name+"-oth", c.otherAllow)
|
||||
|
||||
polMain, merr := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-pergroup-" + c.name + "-main",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpMain.Id},
|
||||
DestinationProviderIds: []string{prov.Id},
|
||||
GuardrailIds: &[]string{gSel.Id},
|
||||
})
|
||||
require.NoError(t, merr, "create main policy %s", c.name)
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMain.Id) })
|
||||
|
||||
polOther, oerr := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-pergroup-" + c.name + "-other",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{grpOther.Id},
|
||||
DestinationProviderIds: []string{prov.Id},
|
||||
GuardrailIds: &[]string{gOth.Id},
|
||||
})
|
||||
require.NoError(t, oerr, "create other policy %s", c.name)
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polOther.Id) })
|
||||
}
|
||||
|
||||
settings, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "read settings")
|
||||
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
|
||||
|
||||
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-pergroup-proxy")
|
||||
require.NoError(t, err, "mint proxy token")
|
||||
px, err := harness.StartProxy(ctx, srv, proxyToken)
|
||||
require.NoError(t, err, "start proxy")
|
||||
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
|
||||
|
||||
cl, err := harness.StartClient(ctx, srv, sk.Key)
|
||||
require.NoError(t, err, "start client")
|
||||
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
|
||||
|
||||
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
|
||||
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
|
||||
require.NoError(t, err, "resolve endpoint to proxy IP")
|
||||
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
|
||||
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
|
||||
}
|
||||
|
||||
send := func(c *pergroupCase, model string) (int, string) {
|
||||
var code int
|
||||
var body string
|
||||
var cerr error
|
||||
switch c.wire {
|
||||
case "vertex":
|
||||
code, body, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, vertexProject, vertexRegion, model, "Reply with exactly: pong", "")
|
||||
case "bedrock":
|
||||
code, body, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, model, "Reply with exactly: pong", "")
|
||||
default:
|
||||
code, body, cerr = cl.Chat(ctx, settings.Endpoint, proxyIP, c.wire, model, "Reply with exactly: pong", "")
|
||||
}
|
||||
require.NoError(t, cerr, "request must reach the proxy for %s", c.name)
|
||||
return code, body
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
// grpMain's own model is served. Retry to absorb tunnel/DNS jitter on
|
||||
// the first call over the freshly warmed tunnel.
|
||||
var code int
|
||||
var body string
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
code, body = send(c, c.selectedReq)
|
||||
if code == 200 {
|
||||
break
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
assert.Equal(t, 200, code,
|
||||
"%s: grpMain's allowlisted model must be served; body: %s\n=== proxy logs ===\n%s", c.name, body, px.Logs(context.Background()))
|
||||
|
||||
// grpOther's model must NOT leak to the grpMain client.
|
||||
code, body = send(c, c.otherReq)
|
||||
assert.Equal(t, 403, code,
|
||||
"%s: another group's allowlisted model must be denied for this caller; body: %s\n=== proxy logs ===\n%s", c.name, body, px.Logs(context.Background()))
|
||||
assert.Contains(t, body, "llm_policy.model_blocked",
|
||||
"%s: denial must be a model-allowlist decision, not routing; body: %s", c.name, body)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardrailMultiGroupUser proves the per-policy/group decision for a caller
|
||||
// that belongs to MULTIPLE groups at once. Two scenarios, one shared stack:
|
||||
//
|
||||
// - union across the user's groups: the client is in gUX and gUY, each with
|
||||
// its own policy+guardrail on provider P1 (gUX->union-a, gUY->union-b). The
|
||||
// client may use BOTH models (the union of its groups' allowlists) while a
|
||||
// third, un-allowlisted model is denied.
|
||||
// - an un-guardrailed group lifts the restriction: the client is in gMP and
|
||||
// gMQ on provider P2, where gMP restricts to mix-a but gMQ's policy carries
|
||||
// NO guardrail. Because one applicable policy is unrestricted, the client may
|
||||
// use a model on no allowlist (mix-z) as well as mix-a.
|
||||
func TestGuardrailMultiGroupUser(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
const (
|
||||
unionA = "mg-union-a"
|
||||
unionB = "mg-union-b"
|
||||
unionC = "mg-union-c" // allowlisted by neither group
|
||||
mixA = "mg-mix-a"
|
||||
mixZ = "mg-mix-z" // on no allowlist; reachable only via the un-guardrailed policy
|
||||
)
|
||||
|
||||
priced := func(ids ...string) *[]api.AgentNetworkProviderModel {
|
||||
out := make([]api.AgentNetworkProviderModel, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
out = append(out, api.AgentNetworkProviderModel{Id: id, InputPer1k: 0.001, OutputPer1k: 0.001})
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
vllm, err := harness.StartVLLM(ctx, srv)
|
||||
require.NoError(t, err, "start mock upstream")
|
||||
t.Cleanup(func() { _ = vllm.Terminate(context.Background()) })
|
||||
|
||||
mkGroup := func(name string) *api.Group {
|
||||
g, gerr := srv.API().Groups.Create(ctx, api.PostApiGroupsJSONRequestBody{Name: name})
|
||||
require.NoError(t, gerr, "create group %s", name)
|
||||
t.Cleanup(func() { _ = srv.API().Groups.Delete(context.Background(), g.Id) })
|
||||
return g
|
||||
}
|
||||
gUX := mkGroup("e2e-mg-union-x")
|
||||
gUY := mkGroup("e2e-mg-union-y")
|
||||
gMP := mkGroup("e2e-mg-mix-p")
|
||||
gMQ := mkGroup("e2e-mg-mix-q")
|
||||
|
||||
ephemeral := false
|
||||
sk, err := srv.API().SetupKeys.Create(ctx, api.PostApiSetupKeysJSONRequestBody{
|
||||
Name: "e2e-mg-client",
|
||||
Type: "reusable",
|
||||
ExpiresIn: 86400,
|
||||
UsageLimit: 0,
|
||||
AutoGroups: []string{gUX.Id, gUY.Id, gMP.Id, gMQ.Id}, // client in all four groups
|
||||
Ephemeral: &ephemeral,
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
require.NotEmpty(t, sk.Key, "setup key plaintext")
|
||||
|
||||
staticKey := "static-e2e-token"
|
||||
enabled := true
|
||||
|
||||
// P1 — union scenario: two restricting policies, one per group.
|
||||
p1, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "e2e-mg-union",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: priced(unionA, unionB, unionC),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
})
|
||||
require.NoError(t, err, "create union provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p1.Id) })
|
||||
|
||||
polUX, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-mg-union-x",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{gUX.Id},
|
||||
DestinationProviderIds: []string{p1.Id},
|
||||
GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-union-x", unionA).Id},
|
||||
})
|
||||
require.NoError(t, err, "create union policy X")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polUX.Id) })
|
||||
|
||||
polUY, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-mg-union-y",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{gUY.Id},
|
||||
DestinationProviderIds: []string{p1.Id},
|
||||
GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-union-y", unionB).Id},
|
||||
})
|
||||
require.NoError(t, err, "create union policy Y")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polUY.Id) })
|
||||
|
||||
// P2 — mixed scenario: one restricting policy + one un-guardrailed policy.
|
||||
p2, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "e2e-mg-mix",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: priced(mixA, mixZ),
|
||||
})
|
||||
require.NoError(t, err, "create mix provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p2.Id) })
|
||||
|
||||
polMP, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-mg-mix-p",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{gMP.Id},
|
||||
DestinationProviderIds: []string{p2.Id},
|
||||
GuardrailIds: &[]string{mkAllowGuardrail(t, ctx, "e2e-mg-mix-p", mixA).Id},
|
||||
})
|
||||
require.NoError(t, err, "create mix policy P")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMP.Id) })
|
||||
|
||||
polMQ, err := srv.CreatePolicy(ctx, api.AgentNetworkPolicyRequest{
|
||||
Name: "e2e-mg-mix-q",
|
||||
Enabled: &enabled,
|
||||
SourceGroups: []string{gMQ.Id},
|
||||
DestinationProviderIds: []string{p2.Id}, // NO guardrail -> unrestricted
|
||||
})
|
||||
require.NoError(t, err, "create mix policy Q")
|
||||
t.Cleanup(func() { _ = srv.DeletePolicy(context.Background(), polMQ.Id) })
|
||||
|
||||
settings, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "read settings")
|
||||
require.NotEmpty(t, settings.Endpoint, "endpoint must be assigned")
|
||||
|
||||
proxyToken, err := srv.CreateProxyTokenCLI(ctx, "e2e-mg-proxy")
|
||||
require.NoError(t, err, "mint proxy token")
|
||||
px, err := harness.StartProxy(ctx, srv, proxyToken)
|
||||
require.NoError(t, err, "start proxy")
|
||||
t.Cleanup(func() { _ = px.Terminate(context.Background()) })
|
||||
|
||||
cl, err := harness.StartClient(ctx, srv, sk.Key)
|
||||
require.NoError(t, err, "start client")
|
||||
t.Cleanup(func() { _ = cl.Terminate(context.Background()) })
|
||||
|
||||
require.NoError(t, cl.WaitConnected(ctx, 90*time.Second), "client must connect to management")
|
||||
proxyIP, err := cl.ResolveProxyIP(ctx, settings.Endpoint)
|
||||
require.NoError(t, err, "resolve endpoint to proxy IP")
|
||||
if err := cl.WaitProxyPeer(ctx, 180*time.Second); err != nil {
|
||||
t.Fatalf("client did not see the proxy peer: %v\n=== proxy logs ===\n%s", err, px.Logs(context.Background()))
|
||||
}
|
||||
|
||||
send := func(model string) (int, string) {
|
||||
code, body, cerr := cl.Chat(ctx, settings.Endpoint, proxyIP, harness.WireChat, model, "Reply with exactly: pong", "")
|
||||
require.NoError(t, cerr, "request must reach the proxy")
|
||||
return code, body
|
||||
}
|
||||
sendUntil200 := func(model string) (int, string) {
|
||||
var code int
|
||||
var body string
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
code, body = send(model)
|
||||
if code == 200 {
|
||||
break
|
||||
}
|
||||
time.Sleep(5 * time.Second)
|
||||
}
|
||||
return code, body
|
||||
}
|
||||
|
||||
t.Run("union across the user's groups", func(t *testing.T) {
|
||||
code, body := sendUntil200(unionA)
|
||||
assert.Equal(t, 200, code, "model allowed by group X must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
|
||||
code, body = sendUntil200(unionB)
|
||||
assert.Equal(t, 200, code, "model allowed by group Y must also be served (union across the user's groups); body: %s", body)
|
||||
|
||||
code, body = send(unionC)
|
||||
assert.Equal(t, 403, code, "a model on neither group's allowlist must be denied; body: %s", body)
|
||||
assert.Contains(t, body, "llm_policy.model_blocked")
|
||||
})
|
||||
|
||||
t.Run("an un-guardrailed group lifts the restriction", func(t *testing.T) {
|
||||
code, body := sendUntil200(mixA)
|
||||
assert.Equal(t, 200, code, "the restricted group's model must be served; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
|
||||
code, body = sendUntil200(mixZ)
|
||||
assert.Equal(t, 200, code,
|
||||
"a non-allowlisted model must be served because the user is also in a group whose policy has no guardrail; body: %s\n=== proxy logs ===\n%s", body, px.Logs(context.Background()))
|
||||
})
|
||||
}
|
||||
|
||||
// mkAllowGuardrail creates a guardrail whose model allowlist is enabled and holds
|
||||
// exactly the given model, registering cleanup.
|
||||
func mkAllowGuardrail(t *testing.T, ctx context.Context, name, model string) api.AgentNetworkGuardrail {
|
||||
t.Helper()
|
||||
var gr api.AgentNetworkGuardrailRequest
|
||||
gr.Name = name
|
||||
gr.Checks.ModelAllowlist.Enabled = true
|
||||
gr.Checks.ModelAllowlist.Models = []string{model}
|
||||
g, err := srv.CreateGuardrail(ctx, gr)
|
||||
require.NoError(t, err, "create guardrail %s", name)
|
||||
t.Cleanup(func() { _ = srv.DeleteGuardrail(context.Background(), g.Id) })
|
||||
return g
|
||||
}
|
||||
@@ -38,7 +38,11 @@ type Proxy struct {
|
||||
// network, registered via the given account proxy token and serving the
|
||||
// AgentNetworkCluster over a self-signed wildcard cert. It does not wait for
|
||||
// peer connectivity — callers poll management for the proxy peer.
|
||||
func StartProxy(ctx context.Context, c *Combined, proxyToken string) (*Proxy, error) {
|
||||
// StartProxy launches the reverse-proxy container. Optional envOverrides are
|
||||
// merged into the container environment after the defaults, so callers can set
|
||||
// or override any NB_PROXY_* var (e.g. NB_PROXY_TUNNEL_CACHE_TTL for tests that
|
||||
// need a short authorization-cache window).
|
||||
func StartProxy(ctx context.Context, c *Combined, proxyToken string, envOverrides ...map[string]string) (*Proxy, error) {
|
||||
root, err := repoRoot()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -93,6 +97,12 @@ func StartProxy(ctx context.Context, c *Combined, proxyToken string) (*Proxy, er
|
||||
WaitingFor: wait.ForLog("Initial mapping sync complete").WithStartupTimeout(90 * time.Second),
|
||||
}
|
||||
|
||||
for _, ov := range envOverrides {
|
||||
for k, v := range ov {
|
||||
req.Env[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
ctr, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{
|
||||
ContainerRequest: req,
|
||||
Started: true,
|
||||
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
|
||||
"github.com/netbirdio/netbird/formatter/hook"
|
||||
agentnetworkpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/server"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
nbdomain "github.com/netbirdio/netbird/shared/management/domain"
|
||||
@@ -112,6 +113,23 @@ var (
|
||||
mgmtSingleAccModeDomain = ""
|
||||
}
|
||||
|
||||
// Load the management-side LLM pricing defaults file: an
|
||||
// explicitly configured path is required to load (a typo must
|
||||
// fail startup — the operator believes those rates are live);
|
||||
// otherwise <datadir>/defaults_llm_pricing.yaml is probed and
|
||||
// may be absent (compiled-in defaults serve). Either way the
|
||||
// path stays watched: the reloader picks up edits — and the
|
||||
// file appearing later — without a restart.
|
||||
pricingPath := config.AgentNetwork.PricingDefaultsFile
|
||||
pricingRequired := pricingPath != ""
|
||||
if !pricingRequired {
|
||||
pricingPath = filepath.Join(config.Datadir, agentnetworkpricing.DefaultFileName)
|
||||
}
|
||||
if err := agentnetworkpricing.LoadFile(pricingPath, pricingRequired); err != nil {
|
||||
return fmt.Errorf("load agent-network pricing defaults: %v", err)
|
||||
}
|
||||
agentnetworkpricing.StartReloader(ctx, agentnetworkpricing.ReloadInterval)
|
||||
|
||||
srv := newServer(&server.Config{
|
||||
NbConfig: config,
|
||||
DNSDomain: dnsDomain,
|
||||
|
||||
@@ -7,12 +7,28 @@ package catalog
|
||||
import "github.com/netbirdio/netbird/shared/management/http/api"
|
||||
|
||||
// Model is the in-memory representation of a catalog model.
|
||||
//
|
||||
// The three cache rates mirror the proxy cost meter's Entry semantics
|
||||
// (USD per 1k tokens; 0 = no rate configured, that bucket bills at
|
||||
// InputPer1k):
|
||||
// - CachedInputPer1k: OpenAI-shape rate for cached prompt tokens
|
||||
// (a SUBSET of input tokens). Typically 0.1-0.5x input.
|
||||
// - CacheReadPer1k / CacheCreationPer1k: Anthropic-shape rates for
|
||||
// the two ADDITIVE prompt-cache buckets. Typically 0.1x / 1.25x
|
||||
// input.
|
||||
//
|
||||
// The catalog is the single default-pricing source: the agentnetwork
|
||||
// pricing package folds these models into per-surface tables that the
|
||||
// synthesizer ships to the proxy's cost_meter.
|
||||
type Model struct {
|
||||
ID string
|
||||
Label string
|
||||
InputPer1k float64
|
||||
OutputPer1k float64
|
||||
ContextWindow int
|
||||
ID string
|
||||
Label string
|
||||
InputPer1k float64
|
||||
OutputPer1k float64
|
||||
CachedInputPer1k float64
|
||||
CacheReadPer1k float64
|
||||
CacheCreationPer1k float64
|
||||
ContextWindow int
|
||||
}
|
||||
|
||||
// ProviderKind groups catalog entries for UI presentation. The split
|
||||
@@ -65,6 +81,17 @@ type Provider struct {
|
||||
// surface — the proxy middleware then falls back to URL sniffing
|
||||
// or skips request-side enrichment.
|
||||
ParserID string
|
||||
// PricingSurfaces names the cost-meter pricing surfaces this
|
||||
// provider's Models are priced under ("openai", "anthropic",
|
||||
// "bedrock" — the llm.Parser surface the request parser stamps as
|
||||
// llm.provider at billing time). NOT derivable from ParserID:
|
||||
// bedrock_api and vertex_ai_api leave ParserID empty (URL-sniffed)
|
||||
// yet price under "bedrock" / "anthropic", and kimi_api serves two
|
||||
// body shapes so it prices under both. Nil for gateway/custom
|
||||
// entries, which declare no models. Same (surface, model) pair
|
||||
// contributed by two providers must carry identical rates — the
|
||||
// pricing package's tests enforce that.
|
||||
PricingSurfaces []string
|
||||
// IdentityInjection, when non-nil, instructs the proxy to stamp
|
||||
// the caller's NetBird identity onto upstream requests under the
|
||||
// configured header names. Used for gateways like LiteLLM that
|
||||
@@ -219,6 +246,7 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#10A37F",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Pricing + context windows cross-checked against LiteLLM's
|
||||
// model_prices_and_context_window.json. Notable corrections from
|
||||
// earlier values: o4-mini repriced from $4/$16 to $1.10/$4.40
|
||||
@@ -226,20 +254,20 @@ var providers = []Provider{
|
||||
// family context windows split between 1.05M for full-size
|
||||
// models and 272K for mini/nano/codex variants.
|
||||
Models: []Model{
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, ContextWindow: 128000},
|
||||
{ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000},
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.5-pro", Label: "GPT-5.5 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-pro", Label: "GPT-5.4 Pro", InputPer1k: 0.030, OutputPer1k: 0.180, CachedInputPer1k: 0.003, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-codex", Label: "GPT-5.3 Codex", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 272000},
|
||||
{ID: "gpt-5.3-chat-latest", Label: "GPT-5.3 Chat", InputPer1k: 0.00175, OutputPer1k: 0.014, CachedInputPer1k: 0.000175, ContextWindow: 128000},
|
||||
{ID: "o4-mini", Label: "o4-mini", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-nano", Label: "GPT-4.1 nano", InputPer1k: 0.0001, OutputPer1k: 0.0004, CachedInputPer1k: 0.000025, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000},
|
||||
{ID: "gpt-4-turbo", Label: "GPT-4 Turbo", InputPer1k: 0.01, OutputPer1k: 0.03, ContextWindow: 128000},
|
||||
{ID: "gpt-3.5-turbo", Label: "GPT-3.5 Turbo", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385},
|
||||
{ID: "text-embedding-3-large", Label: "text-embedding-3-large", InputPer1k: 0.00013, OutputPer1k: 0, ContextWindow: 8191},
|
||||
@@ -257,6 +285,7 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#D97757",
|
||||
ParserID: "anthropic",
|
||||
PricingSurfaces: []string{"anthropic"},
|
||||
// Per Anthropic's current model lineup. Pricing in USD per 1k
|
||||
// tokens. Context windows: 4.6+ family is 1M; Haiku 4.5 stays at
|
||||
// 200K. claude-3-7-sonnet and claude-3-5-haiku retired
|
||||
@@ -267,14 +296,14 @@ var providers = []Provider{
|
||||
// account to be on >= 30-day data retention or all requests
|
||||
// 400.
|
||||
Models: []Model{
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000},
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (deprecated, retires 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000},
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -288,18 +317,19 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#0078D4",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Mirrors openai_api pricing — Azure resells OpenAI models at the
|
||||
// same per-token rates, just under different deployment names.
|
||||
Models: []Model{
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, ContextWindow: 272000},
|
||||
{ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, ContextWindow: 128000},
|
||||
{ID: "gpt-5.5", Label: "GPT-5.5 (Azure)", InputPer1k: 0.005, OutputPer1k: 0.030, CachedInputPer1k: 0.0005, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4", Label: "GPT-5.4 (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.015, CachedInputPer1k: 0.00025, ContextWindow: 1050000},
|
||||
{ID: "gpt-5.4-mini", Label: "GPT-5.4 Mini (Azure)", InputPer1k: 0.00075, OutputPer1k: 0.0045, CachedInputPer1k: 0.000075, ContextWindow: 272000},
|
||||
{ID: "gpt-5.4-nano", Label: "GPT-5.4 Nano (Azure)", InputPer1k: 0.0002, OutputPer1k: 0.00125, CachedInputPer1k: 0.00002, ContextWindow: 272000},
|
||||
{ID: "o4-mini", Label: "o4-mini (Azure)", InputPer1k: 0.0011, OutputPer1k: 0.0044, CachedInputPer1k: 0.000275, ContextWindow: 200000},
|
||||
{ID: "gpt-4.1", Label: "GPT-4.1 (Azure)", InputPer1k: 0.002, OutputPer1k: 0.008, CachedInputPer1k: 0.0005, ContextWindow: 1047576},
|
||||
{ID: "gpt-4.1-mini", Label: "GPT-4.1 mini (Azure)", InputPer1k: 0.0004, OutputPer1k: 0.0016, CachedInputPer1k: 0.0001, ContextWindow: 1047576},
|
||||
{ID: "gpt-4o", Label: "GPT-4o (Azure)", InputPer1k: 0.0025, OutputPer1k: 0.010, CachedInputPer1k: 0.00125, ContextWindow: 128000},
|
||||
{ID: "gpt-4o-mini", Label: "GPT-4o mini (Azure)", InputPer1k: 0.00015, OutputPer1k: 0.0006, CachedInputPer1k: 0.000075, ContextWindow: 128000},
|
||||
{ID: "gpt-35-turbo", Label: "GPT-3.5 Turbo (Azure)", InputPer1k: 0.0005, OutputPer1k: 0.0015, ContextWindow: 16385},
|
||||
},
|
||||
},
|
||||
@@ -313,6 +343,9 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#FF9900",
|
||||
// ParserID stays empty (path-style dispatch via IsBedrockPathStyle);
|
||||
// the request parser meters these under the "bedrock" surface.
|
||||
PricingSurfaces: []string{"bedrock"},
|
||||
// Anthropic models on Bedrock take the anthropic.* prefix and
|
||||
// follow the same lineup / pricing as the first-party Anthropic
|
||||
// catalog entry above. claude-3-7-sonnet and claude-3-5-haiku
|
||||
@@ -322,13 +355,13 @@ var providers = []Provider{
|
||||
// Llama 3.3 70B entry kept unchanged — LiteLLM tracks only
|
||||
// per-region Llama 3 entries; standalone 3.3 not yet listed.
|
||||
Models: []Model{
|
||||
{ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-1", Label: "Claude Opus 4.1 (Bedrock, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000},
|
||||
{ID: "anthropic.claude-haiku-4-5", Label: "Claude Haiku 4.5 (Bedrock)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000},
|
||||
{ID: "meta.llama3-3-70b-instruct", Label: "Llama 3.3 70B (Bedrock)", InputPer1k: 0.00072, OutputPer1k: 0.00072, ContextWindow: 128000},
|
||||
{ID: "amazon.nova-2-lite", Label: "Amazon Nova 2 Lite (Bedrock, preview)", InputPer1k: 0.0003, OutputPer1k: 0.0025, ContextWindow: 1000000},
|
||||
{ID: "amazon.nova-pro", Label: "Amazon Nova Pro (Bedrock)", InputPer1k: 0.0008, OutputPer1k: 0.0032, ContextWindow: 300000},
|
||||
@@ -358,6 +391,10 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#4285F4",
|
||||
// ParserID stays empty (path-style dispatch via IsVertexPathStyle);
|
||||
// Anthropic-on-Vertex requests are metered under the "anthropic"
|
||||
// surface with the bare, unversioned model id.
|
||||
PricingSurfaces: []string{"anthropic"},
|
||||
// Vertex carries the model in the URL path and authenticates with a
|
||||
// service-account-minted OAuth token (api_key = "keyfile::<base64 SA>").
|
||||
// Only Anthropic-on-Vertex is metered today: the request parser maps the
|
||||
@@ -369,14 +406,14 @@ var providers = []Provider{
|
||||
// exists — the router denies unmeterable publishers rather than forward
|
||||
// them uncounted.
|
||||
Models: []Model{
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, ContextWindow: 200000},
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-6", Label: "Claude Opus 4.6 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-1", Label: "Claude Opus 4.1 (Vertex, deprecated 2026-08-05)", InputPer1k: 0.015, OutputPer1k: 0.075, CacheReadPer1k: 0.0015, CacheCreationPer1k: 0.01875, ContextWindow: 200000},
|
||||
{ID: "claude-sonnet-4-6", Label: "Claude Sonnet 4.6 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-4-5", Label: "Claude Sonnet 4.5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 200000},
|
||||
{ID: "claude-haiku-4-5", Label: "Claude Haiku 4.5 (Vertex)", InputPer1k: 0.001, OutputPer1k: 0.005, CacheReadPer1k: 0.0001, CacheCreationPer1k: 0.00125, ContextWindow: 200000},
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -390,6 +427,7 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#FF7000",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Pricing + context windows cross-checked against LiteLLM. Key
|
||||
// gotchas the marketing page hides:
|
||||
// - `mistral-medium-latest` aliases to Medium 3.1 ($0.40/$2),
|
||||
@@ -448,6 +486,10 @@ var providers = []Provider{
|
||||
// model id "k3") is account-bound seat licensing rather than a
|
||||
// meterable platform key, so it's deliberately not the default.
|
||||
ParserID: "",
|
||||
// Both body shapes are metered: /v1/chat/completions under
|
||||
// "openai", /anthropic/v1/messages under "anthropic" — so the
|
||||
// K3 entry is priced on both surfaces.
|
||||
PricingSurfaces: []string{"openai", "anthropic"},
|
||||
// Pricing per Moonshot's platform rates at K3 launch (July 2026):
|
||||
// $3/$15 per MTok with $0.30 cached input, flat across the 1M-token
|
||||
// window. kimi-k3 is the ONLY model the platform serves newer
|
||||
@@ -458,7 +500,14 @@ var providers = []Provider{
|
||||
// The consumer app's "K3 Swarm Max" mode is not an API SKU, so it
|
||||
// doesn't appear here.
|
||||
Models: []Model{
|
||||
{ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, ContextWindow: 1000000},
|
||||
// Carries both cache shapes: Moonshot reports cache hits
|
||||
// OpenAI-style on /v1/chat/completions (CachedInputPer1k)
|
||||
// and Anthropic-style on /anthropic/v1/messages
|
||||
// (CacheReadPer1k) — $0.30/MTok either way. Each surface's
|
||||
// cost formula reads only its own field, so the superset
|
||||
// entry prices both endpoints correctly. No cache-creation
|
||||
// rate published; writes bill at the input rate.
|
||||
{ID: "kimi-k3", Label: "Kimi K3", InputPer1k: 0.003, OutputPer1k: 0.015, CachedInputPer1k: 0.0003, CacheReadPer1k: 0.0003, ContextWindow: 1000000},
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -758,13 +807,28 @@ func IsBedrockPathStyle(providerID string) bool {
|
||||
func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider {
|
||||
models := make([]api.AgentNetworkCatalogModel, 0, len(p.Models))
|
||||
for _, m := range p.Models {
|
||||
models = append(models, api.AgentNetworkCatalogModel{
|
||||
am := api.AgentNetworkCatalogModel{
|
||||
Id: m.ID,
|
||||
Label: m.Label,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
ContextWindow: m.ContextWindow,
|
||||
})
|
||||
}
|
||||
// Cache rates are emitted only when configured so the dashboard
|
||||
// can prefill them; 0 stays off the wire (absent = no rate).
|
||||
if m.CachedInputPer1k > 0 {
|
||||
v := m.CachedInputPer1k
|
||||
am.CachedInputPer1k = &v
|
||||
}
|
||||
if m.CacheReadPer1k > 0 {
|
||||
v := m.CacheReadPer1k
|
||||
am.CacheReadPer1k = &v
|
||||
}
|
||||
if m.CacheCreationPer1k > 0 {
|
||||
v := m.CacheCreationPer1k
|
||||
am.CacheCreationPer1k = &v
|
||||
}
|
||||
models = append(models, am)
|
||||
}
|
||||
kind := api.AgentNetworkCatalogProviderKindProvider
|
||||
switch p.Kind {
|
||||
@@ -784,6 +848,10 @@ func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider {
|
||||
BrandColor: p.BrandColor,
|
||||
Models: models,
|
||||
}
|
||||
if len(p.PricingSurfaces) > 0 {
|
||||
surfaces := append([]string(nil), p.PricingSurfaces...)
|
||||
resp.PricingSurfaces = &surfaces
|
||||
}
|
||||
if len(p.ExtraHeaders) > 0 {
|
||||
extras := make([]api.AgentNetworkCatalogExtraHeader, 0, len(p.ExtraHeaders))
|
||||
for _, h := range p.ExtraHeaders {
|
||||
|
||||
@@ -7,6 +7,7 @@ package handlers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
@@ -15,6 +16,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
@@ -52,11 +54,45 @@ func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) {
|
||||
entries := catalog.All()
|
||||
out := make([]api.AgentNetworkCatalogProvider, 0, len(entries))
|
||||
for _, e := range entries {
|
||||
out = append(out, e.ToAPIResponse())
|
||||
resp := e.ToAPIResponse()
|
||||
applyDefaultPricing(e, &resp)
|
||||
out = append(out, resp)
|
||||
}
|
||||
util.WriteJSONObject(r.Context(), w, out)
|
||||
}
|
||||
|
||||
// applyDefaultPricing overwrites the catalog response's model rates with
|
||||
// the LIVE default pricing table, which may differ from the compiled-in
|
||||
// catalog rates when the operator provides a defaults_llm_pricing.yaml.
|
||||
// This keeps the dashboard's model-row prefill identical to what the
|
||||
// proxy will actually bill — the same table the synthesizer ships.
|
||||
func applyDefaultPricing(cp catalog.Provider, resp *api.AgentNetworkCatalogProvider) {
|
||||
if len(cp.PricingSurfaces) == 0 {
|
||||
return
|
||||
}
|
||||
for i := range resp.Models {
|
||||
m := &resp.Models[i]
|
||||
e, ok := pricing.LookupDefault(cp.PricingSurfaces, m.Id)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
m.InputPer1k = e.InputPer1k
|
||||
m.OutputPer1k = e.OutputPer1k
|
||||
m.CachedInputPer1k = positiveRatePtr(e.CachedInputPer1k)
|
||||
m.CacheReadPer1k = positiveRatePtr(e.CacheReadPer1k)
|
||||
m.CacheCreationPer1k = positiveRatePtr(e.CacheCreationPer1k)
|
||||
}
|
||||
}
|
||||
|
||||
// positiveRatePtr renders a cache rate for the API: absent (nil) when
|
||||
// unset, matching the catalog response convention.
|
||||
func positiveRatePtr(v float64) *float64 {
|
||||
if v <= 0 {
|
||||
return nil
|
||||
}
|
||||
return &v
|
||||
}
|
||||
|
||||
func (h *handler) getAllProviders(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
@@ -213,5 +249,38 @@ func validate(req *api.AgentNetworkProviderRequest, requireAPIKey bool) error {
|
||||
if requireAPIKey && (req.ApiKey == nil || strings.TrimSpace(*req.ApiKey) == "") {
|
||||
return status.Errorf(status.InvalidArgument, "api_key is required")
|
||||
}
|
||||
if req.Models != nil {
|
||||
for i, m := range *req.Models {
|
||||
if err := validateModel(i, m); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateModel is the single ingress guard for operator-entered pricing:
|
||||
// these rates are synthesized into the proxy's cost_meter config verbatim,
|
||||
// and a negative or non-finite rate there would poison every cost the
|
||||
// proxy records, so reject at the API boundary.
|
||||
func validateModel(i int, m api.AgentNetworkProviderModel) error {
|
||||
if strings.TrimSpace(m.Id) == "" {
|
||||
return status.Errorf(status.InvalidArgument, "models[%d]: id is required", i)
|
||||
}
|
||||
rates := map[string]*float64{
|
||||
"input_per_1k": &m.InputPer1k,
|
||||
"output_per_1k": &m.OutputPer1k,
|
||||
"cached_input_per_1k": m.CachedInputPer1k,
|
||||
"cache_read_per_1k": m.CacheReadPer1k,
|
||||
"cache_creation_per_1k": m.CacheCreationPer1k,
|
||||
}
|
||||
for field, v := range rates {
|
||||
if v == nil {
|
||||
continue
|
||||
}
|
||||
if *v < 0 || math.IsNaN(*v) || math.IsInf(*v, 0) {
|
||||
return status.Errorf(status.InvalidArgument, "models[%d] (%s): %s must be a finite, non-negative USD rate", i, m.Id, field)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
func f(v float64) *float64 { return &v }
|
||||
|
||||
// TestValidate_ModelRates guards the single ingress point for operator-entered
|
||||
// pricing. These rates flow verbatim into the proxy's cost_meter config at
|
||||
// synthesis time; the proxy treats a bad rate as a chain-build failure, so
|
||||
// rejecting here is what keeps an account's gateway from going down.
|
||||
func TestValidate_ModelRates(t *testing.T) {
|
||||
base := func(models ...api.AgentNetworkProviderModel) *api.AgentNetworkProviderRequest {
|
||||
key := "sk-test"
|
||||
return &api.AgentNetworkProviderRequest{
|
||||
ProviderId: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamUrl: "https://api.openai.com",
|
||||
ApiKey: &key,
|
||||
Models: &models,
|
||||
}
|
||||
}
|
||||
|
||||
valid := api.AgentNetworkProviderModel{
|
||||
Id: "gpt-4o", InputPer1k: 0.0025, OutputPer1k: 0.01,
|
||||
CachedInputPer1k: f(0.00125),
|
||||
}
|
||||
require.NoError(t, validate(base(valid), true), "finite non-negative rates must pass")
|
||||
|
||||
zeroRates := api.AgentNetworkProviderModel{Id: "self-hosted-llama", InputPer1k: 0, OutputPer1k: 0}
|
||||
require.NoError(t, validate(base(zeroRates), true), "explicit zero prices are allowed (free / self-hosted models)")
|
||||
|
||||
cases := map[string]api.AgentNetworkProviderModel{
|
||||
"empty id": {Id: " ", InputPer1k: 0.001, OutputPer1k: 0.002},
|
||||
"negative input": {Id: "m", InputPer1k: -0.001, OutputPer1k: 0.002},
|
||||
"negative output": {Id: "m", InputPer1k: 0.001, OutputPer1k: -0.002},
|
||||
"NaN input": {Id: "m", InputPer1k: math.NaN(), OutputPer1k: 0.002},
|
||||
"Inf output": {Id: "m", InputPer1k: 0.001, OutputPer1k: math.Inf(1)},
|
||||
"negative cached": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CachedInputPer1k: f(-1)},
|
||||
"NaN cache read": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheReadPer1k: f(math.NaN())},
|
||||
"Inf cache creation": {Id: "m", InputPer1k: 0.001, OutputPer1k: 0.002, CacheCreationPer1k: f(math.Inf(-1))},
|
||||
}
|
||||
for name, m := range cases {
|
||||
assert.Error(t, validate(base(m), true), "case %q must be rejected", name)
|
||||
}
|
||||
}
|
||||
@@ -86,12 +86,17 @@ type Manager interface {
|
||||
|
||||
// PolicySelectionInput is the per-request selection envelope. The
|
||||
// proxy populates it from CapturedData (account, user, groups) plus
|
||||
// the provider llm_router resolved.
|
||||
// the provider llm_router resolved and the model it extracted.
|
||||
type PolicySelectionInput struct {
|
||||
AccountID string
|
||||
UserID string
|
||||
GroupIDs []string
|
||||
ProviderID string
|
||||
// Model is the already-normalised upstream model id the proxy extracted
|
||||
// (parser strips Bedrock region/version, Vertex @version), so a
|
||||
// case-insensitive compare suffices. Empty = undetermined → not permitted
|
||||
// (fail closed).
|
||||
Model string
|
||||
}
|
||||
|
||||
// PolicySelectionResult names the policy that "pays" for this request
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"math"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
@@ -35,6 +36,10 @@ const (
|
||||
denyCodeAccountTokenCapExceeded = "llm_account.token_cap_exceeded"
|
||||
//nolint:gosec // account deny code label, not a credential
|
||||
denyCodeAccountBudgetCapExceeded = "llm_account.budget_cap_exceeded"
|
||||
// denyCodeModelBlocked is returned when policies govern the request's
|
||||
// (provider, caller-groups) but none permits the model. Matches the proxy
|
||||
// guardrail's code so both layers surface the same label.
|
||||
denyCodeModelBlocked = "llm_policy.model_blocked"
|
||||
)
|
||||
|
||||
// consumptionCache holds the consumption counters prefetched for one
|
||||
@@ -159,6 +164,25 @@ func (m *managerImpl) SelectPolicyForRequest(ctx context.Context, in PolicySelec
|
||||
}
|
||||
candidates := filterApplicablePolicies(policies, in)
|
||||
|
||||
// Model-allowlist gate scoped to the matched policies: keep candidates whose
|
||||
// guardrails permit the model (none enabled = unrestricted), deny when
|
||||
// policies apply but none permits it. Skip the load when none has a guardrail.
|
||||
if len(candidates) > 0 && anyPolicyHasGuardrails(candidates) {
|
||||
guardrailsByID, gErr := m.loadGuardrailsByID(ctx, in.AccountID)
|
||||
if gErr != nil {
|
||||
return nil, gErr
|
||||
}
|
||||
permitted := filterModelPermittedPolicies(candidates, guardrailsByID, in.Model)
|
||||
if len(permitted) == 0 {
|
||||
return &PolicySelectionResult{
|
||||
Allow: false,
|
||||
DenyCode: denyCodeModelBlocked,
|
||||
DenyReason: modelBlockedReason(in.Model),
|
||||
}, nil
|
||||
}
|
||||
candidates = permitted
|
||||
}
|
||||
|
||||
// Prefetch every consumption counter the ceiling + candidate policies will
|
||||
// read, in a single store round-trip, then score against the cache.
|
||||
cache, err := m.prefetchConsumption(ctx, in, rules, candidates, now)
|
||||
@@ -250,6 +274,90 @@ func filterApplicablePolicies(policies []*types.Policy, in PolicySelectionInput)
|
||||
return out
|
||||
}
|
||||
|
||||
// anyPolicyHasGuardrails reports whether any policy references at least one
|
||||
// guardrail, so the selector can skip loading guardrails when none do.
|
||||
func anyPolicyHasGuardrails(policies []*types.Policy) bool {
|
||||
for _, p := range policies {
|
||||
if p != nil && len(p.GuardrailIDs) > 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// loadGuardrailsByID loads the account's guardrails indexed by ID. Used by the
|
||||
// model-allowlist gate to resolve each candidate policy's attached guardrails.
|
||||
func (m *managerImpl) loadGuardrailsByID(ctx context.Context, accountID string) (map[string]*types.Guardrail, error) {
|
||||
guardrails, err := m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list account guardrails: %w", err)
|
||||
}
|
||||
byID := make(map[string]*types.Guardrail, len(guardrails))
|
||||
for _, g := range guardrails {
|
||||
if g != nil {
|
||||
byID[g.ID] = g
|
||||
}
|
||||
}
|
||||
return byID, nil
|
||||
}
|
||||
|
||||
// filterModelPermittedPolicies returns the subset of policies whose guardrails
|
||||
// permit the model. Order is preserved so downstream scoring is unaffected.
|
||||
func filterModelPermittedPolicies(policies []*types.Policy, byID map[string]*types.Guardrail, model string) []*types.Policy {
|
||||
out := make([]*types.Policy, 0, len(policies))
|
||||
for _, p := range policies {
|
||||
if policyPermitsModel(p, byID, model) {
|
||||
out = append(out, p)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// policyPermitsModel reports whether a policy permits the model. No
|
||||
// allowlist-enabled guardrail = unrestricted (permits any, incl. empty);
|
||||
// otherwise the model must be in the union of its allowlists, so an
|
||||
// empty/undetermined model fails closed.
|
||||
func policyPermitsModel(p *types.Policy, byID map[string]*types.Guardrail, model string) bool {
|
||||
if p == nil {
|
||||
return false
|
||||
}
|
||||
wanted := normaliseModelID(model)
|
||||
restricted := false
|
||||
for _, gID := range p.GuardrailIDs {
|
||||
g, ok := byID[gID]
|
||||
if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled {
|
||||
continue
|
||||
}
|
||||
restricted = true
|
||||
if wanted == "" {
|
||||
continue
|
||||
}
|
||||
for _, allowed := range g.Checks.ModelAllowlist.Models {
|
||||
if normaliseModelID(allowed) == wanted {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return !restricted
|
||||
}
|
||||
|
||||
// normaliseModelID lowercases and trims a model identifier so the allowlist
|
||||
// compare is case-insensitive and trim-tolerant. Mirrors the proxy guardrail's
|
||||
// normaliseModel so both layers agree on what "same model" means.
|
||||
func normaliseModelID(model string) string {
|
||||
return strings.ToLower(strings.TrimSpace(model))
|
||||
}
|
||||
|
||||
// modelBlockedReason builds the human-readable deny reason for a model-allowlist
|
||||
// rejection. The model is quoted when known; an undetermined model is reported
|
||||
// as such so the access log distinguishes "wrong model" from "no model".
|
||||
func modelBlockedReason(model string) string {
|
||||
if normaliseModelID(model) == "" {
|
||||
return "request model could not be determined for the policy allowlist"
|
||||
}
|
||||
return fmt.Sprintf("model %q is not permitted by any applicable policy allowlist", model)
|
||||
}
|
||||
|
||||
// candidate is the per-policy intermediate the selector ranks. A
|
||||
// policy that's been exhausted on any enabled cap never makes it
|
||||
// into this slice; the selector's deny envelope carries the latest
|
||||
|
||||
@@ -0,0 +1,329 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
)
|
||||
|
||||
// guardedPolicy builds an enabled, uncapped policy that authorises sourceGroups
|
||||
// to reach providerID under the given guardrails. Uncapped keeps the selector's
|
||||
// headroom scoring trivial so these tests isolate the model-allowlist gate.
|
||||
func guardedPolicy(id, account string, sourceGroups []string, providerID string, guardrailIDs ...string) *types.Policy {
|
||||
return &types.Policy{
|
||||
ID: id,
|
||||
AccountID: account,
|
||||
Enabled: true,
|
||||
SourceGroups: sourceGroups,
|
||||
DestinationProviderIDs: []string{providerID},
|
||||
GuardrailIDs: guardrailIDs,
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}
|
||||
}
|
||||
|
||||
// allowlistGuardrail builds a guardrail whose model allowlist is enabled and
|
||||
// carries the given models.
|
||||
func allowlistGuardrail(id, account string, models ...string) *types.Guardrail {
|
||||
return &types.Guardrail{
|
||||
ID: id,
|
||||
AccountID: account,
|
||||
Checks: types.GuardrailChecks{
|
||||
ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true, Models: models},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func expectPolicies(mockStore *store.MockStore, account string, policies ...*types.Policy) {
|
||||
mockStore.EXPECT().
|
||||
GetAccountAgentNetworkPolicies(gomock.Any(), gomock.Any(), account).
|
||||
Return(policies, nil)
|
||||
}
|
||||
|
||||
func expectGuardrails(mockStore *store.MockStore, account string, guardrails ...*types.Guardrail) {
|
||||
mockStore.EXPECT().
|
||||
GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), account).
|
||||
Return(guardrails, nil)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_ModelBlockedByAllowlist proves the authoritative allowlist
|
||||
// decision: a policy authorises the (provider, group) but restricts the model,
|
||||
// and the requested model isn't on the list, so the request is denied.
|
||||
func TestSelectPolicy_ModelBlockedByAllowlist(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o"))
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
UserID: "user-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "claude-opus-4",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, res.Allow, "a model outside the only applicable policy's allowlist must be denied")
|
||||
assert.Equal(t, denyCodeModelBlocked, res.DenyCode, "deny code must be model_blocked")
|
||||
assert.NotEmpty(t, res.DenyReason, "deny reason must be populated")
|
||||
}
|
||||
|
||||
// TestSelectPolicy_ModelAllowedByAllowlist is the allow counterpart: the model
|
||||
// is on the applicable policy's allowlist, so selection proceeds normally.
|
||||
func TestSelectPolicy_ModelAllowedByAllowlist(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o", "claude-opus-4"))
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
UserID: "user-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "claude-opus-4",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow, "a model on the applicable policy's allowlist must be allowed")
|
||||
assert.Equal(t, "pol-A", res.SelectedPolicyID)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_CaseInsensitiveModelMatch proves the compare tolerates case
|
||||
// and surrounding whitespace, matching the proxy guardrail's normalisation.
|
||||
func TestSelectPolicy_CaseInsensitiveModelMatch(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", " GPT-4o "))
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "gpt-4o",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow, "case/whitespace variants must match the allowlist entry")
|
||||
}
|
||||
|
||||
// TestSelectPolicy_UnguardedPolicyIsUnrestricted is the false-deny fix: when two
|
||||
// policies authorise the same (provider, group) and one has no guardrail, that
|
||||
// policy makes the request unrestricted — not caught by the other's allowlist.
|
||||
func TestSelectPolicy_UnguardedPolicyIsUnrestricted(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
restricted := guardedPolicy("pol-restricted", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
open := guardedPolicy("pol-open", "acc-1", []string{"grp-eng"}, "prov-1") // no guardrail
|
||||
expectPolicies(mockStore, "acc-1", restricted, open)
|
||||
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o"))
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "claude-opus-4",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow, "an un-guardrailed policy for the same (provider, group) must leave the request unrestricted")
|
||||
assert.Equal(t, "pol-open", res.SelectedPolicyID, "the unrestricted policy must be the one that pays")
|
||||
}
|
||||
|
||||
// TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups is the false-allow fix: a
|
||||
// model allowlisted only for grp-b must not be usable by a grp-a caller. The
|
||||
// selector considers only policies applicable to the caller's groups.
|
||||
func TestSelectPolicy_AllowlistDoesNotLeakAcrossGroups(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
polA := guardedPolicy("pol-a", "acc-1", []string{"grp-a"}, "prov-1", "g-a")
|
||||
polB := guardedPolicy("pol-b", "acc-1", []string{"grp-b"}, "prov-1", "g-b")
|
||||
expectPolicies(mockStore, "acc-1", polA, polB)
|
||||
expectGuardrails(mockStore, "acc-1",
|
||||
allowlistGuardrail("g-a", "acc-1", "gpt-4o"),
|
||||
allowlistGuardrail("g-b", "acc-1", "claude-opus-4"),
|
||||
)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-a"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "claude-opus-4", // only allowed for grp-b
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, res.Allow, "grp-b's allowlisted model must not leak to a grp-a caller")
|
||||
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_UndeterminedModelFailsClosed proves the fail-closed contract
|
||||
// mirrors the proxy: with a restricted applicable policy and an empty model
|
||||
// (e.g. a path-routed shape the parser couldn't map), the request is denied.
|
||||
func TestSelectPolicy_UndeterminedModelFailsClosed(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1", allowlistGuardrail("g-1", "acc-1", "gpt-4o"))
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "", // undetermined
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, res.Allow, "an undetermined model must fail closed against a restricted policy")
|
||||
assert.Equal(t, denyCodeModelBlocked, res.DenyCode)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_DisabledAllowlistDoesNotRestrict proves a guardrail whose
|
||||
// model allowlist is disabled imposes no model restriction, even though the
|
||||
// policy references it.
|
||||
func TestSelectPolicy_DisabledAllowlistDoesNotRestrict(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
disabled := &types.Guardrail{
|
||||
ID: "g-1",
|
||||
AccountID: "acc-1",
|
||||
Checks: types.GuardrailChecks{
|
||||
ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}},
|
||||
},
|
||||
}
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1", disabled)
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "anything-goes",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow, "a disabled allowlist must not restrict the model")
|
||||
assert.Equal(t, "pol-A", res.SelectedPolicyID)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_UnionAcrossPolicyGuardrails proves a policy with multiple
|
||||
// allowlist guardrails permits the union of their models (not just the first).
|
||||
func TestSelectPolicy_UnionAcrossPolicyGuardrails(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1", "g-2")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1",
|
||||
allowlistGuardrail("g-1", "acc-1", "gpt-4o"),
|
||||
allowlistGuardrail("g-2", "acc-1", "claude-opus-4"),
|
||||
)
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "claude-opus-4", // only in the second guardrail's list
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow, "a model in any of the policy's allowlist guardrails must be permitted")
|
||||
assert.Equal(t, "pol-A", res.SelectedPolicyID)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_GuardrailLookupErrorPropagates proves a store failure while
|
||||
// resolving the candidate policies' guardrails surfaces as an error, not a
|
||||
// silent allow/deny.
|
||||
func TestSelectPolicy_GuardrailLookupErrorPropagates(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-1")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
mockStore.EXPECT().
|
||||
GetAccountAgentNetworkGuardrails(gomock.Any(), gomock.Any(), "acc-1").
|
||||
Return(nil, errors.New("store unavailable"))
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "gpt-4o",
|
||||
})
|
||||
require.Error(t, err, "a guardrail-lookup failure must surface as an error")
|
||||
assert.Nil(t, res)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted proves a
|
||||
// policy referencing a guardrail ID absent from the account's set (a stale
|
||||
// reference) imposes no model restriction — same as no guardrail.
|
||||
func TestSelectPolicy_MissingGuardrailReferenceTreatedAsUnrestricted(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
policy := guardedPolicy("pol-A", "acc-1", []string{"grp-eng"}, "prov-1", "g-missing")
|
||||
expectPolicies(mockStore, "acc-1", policy)
|
||||
expectGuardrails(mockStore, "acc-1")
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "anything-goes",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow, "an orphaned guardrail reference must not restrict the model")
|
||||
assert.Equal(t, "pol-A", res.SelectedPolicyID)
|
||||
}
|
||||
|
||||
// TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter proves the model
|
||||
// gate narrows candidates before cap scoring: the permitting policy is selected
|
||||
// even though the blocked one has a larger, more attractive cap.
|
||||
func TestSelectPolicy_PartialCandidatesPermittedAfterModelFilter(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr, mockStore := newSelectorMgr(t, ctrl)
|
||||
|
||||
polBig := guardedPolicy("pol-big", "acc-1", []string{"grp-eng"}, "prov-1", "g-restrict")
|
||||
polBig.Limits = types.PolicyLimits{
|
||||
TokenLimit: types.PolicyTokenLimit{Enabled: true, GroupCap: 1_000_000, WindowSeconds: 3600},
|
||||
}
|
||||
polSmall := guardedPolicy("pol-small", "acc-1", []string{"grp-eng"}, "prov-1", "g-permit")
|
||||
polSmall.Limits = types.PolicyLimits{
|
||||
TokenLimit: types.PolicyTokenLimit{Enabled: true, GroupCap: 100, WindowSeconds: 3600},
|
||||
}
|
||||
expectPolicies(mockStore, "acc-1", polBig, polSmall)
|
||||
expectGuardrails(mockStore, "acc-1",
|
||||
allowlistGuardrail("g-restrict", "acc-1", "gpt-4o"),
|
||||
allowlistGuardrail("g-permit", "acc-1", "claude-opus-4"),
|
||||
)
|
||||
expectConsumptionBatch(mockStore, nil)
|
||||
|
||||
res, err := mgr.SelectPolicyForRequest(context.Background(), PolicySelectionInput{
|
||||
AccountID: "acc-1",
|
||||
GroupIDs: []string{"grp-eng"},
|
||||
ProviderID: "prov-1",
|
||||
Model: "claude-opus-4", // only pol-small's guardrail permits this
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, res.Allow)
|
||||
assert.Equal(t, "pol-small", res.SelectedPolicyID,
|
||||
"the model filter must exclude pol-big before cap scoring")
|
||||
}
|
||||
156
management/internals/modules/agentnetwork/pricing/defaults.go
Normal file
156
management/internals/modules/agentnetwork/pricing/defaults.go
Normal file
@@ -0,0 +1,156 @@
|
||||
// Package pricing builds the default LLM pricing table the synthesizer
|
||||
// ships to the proxy's cost_meter middleware. The catalog is the single
|
||||
// source of default rates: every catalog provider's models are folded
|
||||
// into the pricing surfaces the provider declares (PricingSurfaces),
|
||||
// then a small supplemental list adds priced-but-not-operator-selectable
|
||||
// entries. Management is the sole pricing authority — the proxy carries
|
||||
// no embedded price list and bills exclusively from the table it is
|
||||
// sent.
|
||||
//go:generate go run gen.go
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
)
|
||||
|
||||
// Entry is a single model's pricing in USD per 1k tokens. This struct IS
|
||||
// the wire shape: the synthesizer marshals it verbatim into cost_meter's
|
||||
// ConfigJSON, and the proxy unmarshals the same field names.
|
||||
//
|
||||
// A zero rate means "no rate configured" — the proxy bills that cache
|
||||
// bucket at InputPer1k (identical semantics to the retired proxy-embedded
|
||||
// table). CachedInputPer1k is the OpenAI shape (cached prompt tokens are
|
||||
// a subset of input); CacheReadPer1k / CacheCreationPer1k are the
|
||||
// Anthropic shape (additive buckets).
|
||||
type Entry struct {
|
||||
InputPer1k float64 `json:"input_per_1k"`
|
||||
OutputPer1k float64 `json:"output_per_1k"`
|
||||
CachedInputPer1k float64 `json:"cached_input_per_1k,omitempty"`
|
||||
CacheReadPer1k float64 `json:"cache_read_per_1k,omitempty"`
|
||||
CacheCreationPer1k float64 `json:"cache_creation_per_1k,omitempty"`
|
||||
}
|
||||
|
||||
// supplementalDefaults are (surface, model) entries that are priced but
|
||||
// deliberately not operator-selectable in the catalog. Each carries a
|
||||
// reason; when one of these models joins a catalog lineup, delete the
|
||||
// row here — the collision test fails loudly if the rates ever disagree.
|
||||
var supplementalDefaults = map[string]map[string]Entry{
|
||||
"openai": {
|
||||
// GPT-5 (2025) family — kept for gateway requests using the
|
||||
// unsuffixed ids; the dashboard offers only the 5.x lineup.
|
||||
"gpt-5": {InputPer1k: 0.00125, OutputPer1k: 0.01, CachedInputPer1k: 0.000125},
|
||||
"gpt-5-mini": {InputPer1k: 0.00025, OutputPer1k: 0.002, CachedInputPer1k: 0.000025},
|
||||
"gpt-5-nano": {InputPer1k: 0.00005, OutputPer1k: 0.0004, CachedInputPer1k: 0.000005},
|
||||
},
|
||||
"anthropic": {
|
||||
// claude-opus-5 is not yet in the catalog lineup but gateway /
|
||||
// grandfathered traffic uses it; priced so it isn't skipped.
|
||||
"claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
|
||||
// "kimi-k3[1m]" is the 1M-context alias some Claude Code guides
|
||||
// configure against Moonshot's Anthropic-compatible endpoint;
|
||||
// priced identically to kimi-k3 so those requests aren't skipped.
|
||||
"kimi-k3[1m]": {InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003},
|
||||
},
|
||||
"bedrock": {
|
||||
"anthropic.claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
|
||||
},
|
||||
}
|
||||
|
||||
var (
|
||||
compiledOnce sync.Once
|
||||
compiledTable map[string]map[string]Entry
|
||||
// mergedTable holds the current live table when a pricing defaults
|
||||
// file is loaded: the file merged entry-whole over the compiled-in
|
||||
// base. Nil while no file is loaded (or after the file is removed),
|
||||
// in which case the compiled-in table serves. Swapped atomically by
|
||||
// the file loader/reloader; readers never block.
|
||||
mergedTable atomic.Pointer[map[string]map[string]Entry]
|
||||
)
|
||||
|
||||
// DefaultTable returns the current default pricing table keyed
|
||||
// surface -> model -> Entry: the management-side defaults file (see
|
||||
// LoadFile / StartReloader) when one is loaded, merged over the
|
||||
// compiled-in catalog table, which alone serves as the fallback when no
|
||||
// file exists. The snapshot may change between calls as the file is
|
||||
// re-read — consumers (the synthesizer on every reconcile, the catalog
|
||||
// endpoint on every request) pick up fresh rates automatically. Callers
|
||||
// must not mutate the returned maps.
|
||||
func DefaultTable() map[string]map[string]Entry {
|
||||
if t := mergedTable.Load(); t != nil {
|
||||
return *t
|
||||
}
|
||||
return compiledBase()
|
||||
}
|
||||
|
||||
// compiledBase returns the compiled-in table (catalog + supplementals),
|
||||
// built once.
|
||||
func compiledBase() map[string]map[string]Entry {
|
||||
compiledOnce.Do(func() {
|
||||
compiledTable = buildDefaultTable()
|
||||
})
|
||||
return compiledTable
|
||||
}
|
||||
|
||||
func buildDefaultTable() map[string]map[string]Entry {
|
||||
out := make(map[string]map[string]Entry)
|
||||
for _, p := range catalog.All() {
|
||||
for _, surface := range p.PricingSurfaces {
|
||||
inner, ok := out[surface]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry)
|
||||
out[surface] = inner
|
||||
}
|
||||
for _, m := range p.Models {
|
||||
// First writer wins; providers contributing the same
|
||||
// (surface, model) must agree on rates — enforced by
|
||||
// TestDefaultTable_NoConflictingContributions.
|
||||
if _, dup := inner[m.ID]; dup {
|
||||
continue
|
||||
}
|
||||
inner[m.ID] = entryFromCatalogModel(m)
|
||||
}
|
||||
}
|
||||
}
|
||||
for surface, models := range supplementalDefaults {
|
||||
inner, ok := out[surface]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry)
|
||||
out[surface] = inner
|
||||
}
|
||||
for id, e := range models {
|
||||
if _, dup := inner[id]; dup {
|
||||
continue
|
||||
}
|
||||
inner[id] = e
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func entryFromCatalogModel(m catalog.Model) Entry {
|
||||
return Entry{
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
CachedInputPer1k: m.CachedInputPer1k,
|
||||
CacheReadPer1k: m.CacheReadPer1k,
|
||||
CacheCreationPer1k: m.CacheCreationPer1k,
|
||||
}
|
||||
}
|
||||
|
||||
// LookupDefault returns the default entry for model on the first of the
|
||||
// given surfaces that prices it. Used by the synthesizer to seed a
|
||||
// per-provider entry with default cache rates before overlaying the
|
||||
// operator's stored prices.
|
||||
func LookupDefault(surfaces []string, model string) (Entry, bool) {
|
||||
table := DefaultTable()
|
||||
for _, s := range surfaces {
|
||||
if e, ok := table[s][model]; ok {
|
||||
return e, true
|
||||
}
|
||||
}
|
||||
return Entry{}, false
|
||||
}
|
||||
@@ -1,83 +1,176 @@
|
||||
# Embedded default pricing for llm_observability. Compiled into the proxy
|
||||
# binary via go:embed in pricing.go; cost annotation works out of the box
|
||||
# without any operator action.
|
||||
# Default LLM pricing used by NetBird's Agent Network cost metering.
|
||||
# GENERATED from the management catalog — do not edit this file in the
|
||||
# repository; regenerate with:
|
||||
#
|
||||
# Operators override entries by dropping a pricing.yaml into --plugin-data-dir
|
||||
# (or whichever basename is given via params.pricing_path). The override file
|
||||
# only needs entries the operator wants to change; missing entries fall
|
||||
# through to these defaults.
|
||||
# go generate ./management/internals/modules/agentnetwork/pricing
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Public list prices drift; ship a fresh
|
||||
# binary or override individual entries via the override file as needed.
|
||||
# Operators: copy this file to <datadir>/defaults_llm_pricing.yaml (or
|
||||
# any path configured via management.json:
|
||||
#
|
||||
# Optional cache fields:
|
||||
# cached_input_per_1k OpenAI: rate for prompt_tokens_details.cached_tokens
|
||||
# (a SUBSET of prompt_tokens). Typically 0.5x input.
|
||||
# Absent → cached portion bills at input_per_1k.
|
||||
# cache_read_per_1k Anthropic: rate for cache_read_input_tokens
|
||||
# (ADDITIVE to input_tokens). Typically 0.1x input.
|
||||
# Absent → cache reads bill at input_per_1k.
|
||||
# cache_creation_per_1k Anthropic: rate for cache_creation_input_tokens
|
||||
# (ADDITIVE to input_tokens). Typically 1.25x input.
|
||||
# Absent → cache writes bill at input_per_1k.
|
||||
# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } }
|
||||
#
|
||||
# ) and adjust the entries you want to change. Management re-reads the
|
||||
# file periodically (mtime poll, every minute): the live table feeds the
|
||||
# proxies' cost metering and the dashboard's model-price prefill, so
|
||||
# edits apply without a restart. Your file only needs the entries you
|
||||
# want to change — but each entry REPLACES the built-in entry for that
|
||||
# surface+model whole, so repeat the cache rates you want to keep.
|
||||
# Unknown fields and negative or non-finite rates are rejected: at
|
||||
# startup that fails boot (for an explicitly configured path); at
|
||||
# runtime the previous table is kept and a warning is logged. Deleting
|
||||
# the file reverts to the built-in defaults below.
|
||||
#
|
||||
# Top-level keys are pricing surfaces — the parser shape requests are
|
||||
# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible
|
||||
# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock"
|
||||
# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be
|
||||
# the normalized id the proxy meters (version/region suffixes stripped).
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Optional cache fields:
|
||||
# cached_input_per_1k OpenAI shape: rate for cached prompt tokens
|
||||
# (a SUBSET of input tokens). Absent -> cached
|
||||
# portion bills at input_per_1k.
|
||||
# cache_read_per_1k Anthropic shape: rate for cache_read tokens
|
||||
# (ADDITIVE to input). Absent -> input rate.
|
||||
# cache_creation_per_1k Anthropic shape: rate for cache_creation
|
||||
# tokens (ADDITIVE to input). Absent -> input
|
||||
# rate.
|
||||
|
||||
anthropic:
|
||||
claude-fable-5:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.05
|
||||
cache_read_per_1k: 0.001
|
||||
cache_creation_per_1k: 0.0125
|
||||
claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.0003
|
||||
cache_read_per_1k: 0.0003
|
||||
"kimi-k3[1m]":
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
|
||||
bedrock:
|
||||
amazon.nova-2-lite:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0025
|
||||
amazon.nova-lite:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00024
|
||||
amazon.nova-micro:
|
||||
input_per_1k: 0.000035
|
||||
output_per_1k: 0.00014
|
||||
amazon.nova-pro:
|
||||
input_per_1k: 0.0008
|
||||
output_per_1k: 0.0032
|
||||
anthropic.claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
anthropic.claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
anthropic.claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
meta.llama3-3-70b-instruct:
|
||||
input_per_1k: 0.00072
|
||||
output_per_1k: 0.00072
|
||||
|
||||
openai:
|
||||
# OpenAI + OpenAI-compatible providers (openai_api, azure_openai_api,
|
||||
# mistral_api, and the openai-parser gateways) all emit llm.provider="openai",
|
||||
# so their models are priced here. Kept in sync with the management catalog;
|
||||
# rates cross-checked against LiteLLM model_prices_and_context_window.json.
|
||||
|
||||
# GPT-5.x family — cache reads 10% of input (0.1x).
|
||||
gpt-5.5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.03
|
||||
cached_input_per_1k: 0.0005
|
||||
gpt-5.5-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
gpt-5.4:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.00025
|
||||
gpt-5.4-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
gpt-5.4-mini:
|
||||
input_per_1k: 0.00075
|
||||
output_per_1k: 0.0045
|
||||
cached_input_per_1k: 0.000075
|
||||
gpt-5.4-nano:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.00125
|
||||
cached_input_per_1k: 0.00002
|
||||
gpt-5.3-codex:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
gpt-5.3-chat-latest:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
# GPT-5 (2025) family — kept for gateway requests using the unsuffixed ids.
|
||||
gpt-5:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
cached_input_per_1k: 0.000125
|
||||
gpt-5-mini:
|
||||
input_per_1k: 0.00025
|
||||
codestral-2508:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0009
|
||||
codestral-latest:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.003
|
||||
devstral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
cached_input_per_1k: 0.000025
|
||||
gpt-5-nano:
|
||||
input_per_1k: 0.00005
|
||||
output_per_1k: 0.0004
|
||||
cached_input_per_1k: 0.000005
|
||||
o4-mini:
|
||||
input_per_1k: 0.0011
|
||||
output_per_1k: 0.0044
|
||||
cached_input_per_1k: 0.000275
|
||||
# GPT-4.1 family — cache reads 25% of input.
|
||||
devstral-small-latest:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0003
|
||||
gpt-3.5-turbo:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
gpt-35-turbo:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
gpt-4-turbo:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.03
|
||||
gpt-4.1:
|
||||
input_per_1k: 0.002
|
||||
output_per_1k: 0.008
|
||||
@@ -90,7 +183,6 @@ openai:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0004
|
||||
cached_input_per_1k: 0.000025
|
||||
# GPT-4o family — cache reads 50% of input (0.5x).
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
@@ -99,200 +191,92 @@ openai:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
cached_input_per_1k: 0.000075
|
||||
# Older GPT — no prompt caching.
|
||||
gpt-4-turbo:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.03
|
||||
gpt-3.5-turbo:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
gpt-35-turbo: # Azure deployment alias of gpt-3.5-turbo
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
# Embeddings — no caching, no output tokens.
|
||||
text-embedding-3-large:
|
||||
input_per_1k: 0.00013
|
||||
output_per_1k: 0
|
||||
text-embedding-3-small:
|
||||
input_per_1k: 0.00002
|
||||
output_per_1k: 0
|
||||
|
||||
# Mistral (mistral_api) — routed via the openai parser; no prompt caching.
|
||||
mistral-large-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
mistral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
gpt-5:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
cached_input_per_1k: 0.000125
|
||||
gpt-5-mini:
|
||||
input_per_1k: 0.00025
|
||||
output_per_1k: 0.002
|
||||
mistral-medium-3-5:
|
||||
input_per_1k: 0.0015
|
||||
output_per_1k: 0.0075
|
||||
mistral-small-latest:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00018
|
||||
cached_input_per_1k: 0.000025
|
||||
gpt-5-nano:
|
||||
input_per_1k: 0.00005
|
||||
output_per_1k: 0.0004
|
||||
cached_input_per_1k: 0.000005
|
||||
gpt-5.3-chat-latest:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
gpt-5.3-codex:
|
||||
input_per_1k: 0.00175
|
||||
output_per_1k: 0.014
|
||||
cached_input_per_1k: 0.000175
|
||||
gpt-5.4:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.00025
|
||||
gpt-5.4-mini:
|
||||
input_per_1k: 0.00075
|
||||
output_per_1k: 0.0045
|
||||
cached_input_per_1k: 0.000075
|
||||
gpt-5.4-nano:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.00125
|
||||
cached_input_per_1k: 0.00002
|
||||
gpt-5.4-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
gpt-5.5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.03
|
||||
cached_input_per_1k: 0.0005
|
||||
gpt-5.5-pro:
|
||||
input_per_1k: 0.03
|
||||
output_per_1k: 0.18
|
||||
cached_input_per_1k: 0.003
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.0003
|
||||
cache_read_per_1k: 0.0003
|
||||
magistral-medium-latest:
|
||||
input_per_1k: 0.002
|
||||
output_per_1k: 0.005
|
||||
magistral-small-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
devstral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
devstral-small-latest:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0003
|
||||
codestral-2508:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0009
|
||||
codestral-latest:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.003
|
||||
ministral-3-14b-2512:
|
||||
input_per_1k: 0.0002
|
||||
output_per_1k: 0.0002
|
||||
ministral-8b-latest:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.00015
|
||||
ministral-3-3b-2512:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0.0001
|
||||
ministral-8b-latest:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.00015
|
||||
mistral-embed:
|
||||
input_per_1k: 0.0001
|
||||
output_per_1k: 0
|
||||
|
||||
# Kimi / Moonshot AI (kimi_api) — OpenAI-compatible /v1 endpoint. Moonshot
|
||||
# reports cache hits OpenAI-style when present; cached input is 10% of
|
||||
# input ($0.30 vs $3.00 per MTok). kimi-k3 is the only model the platform
|
||||
# serves newer accounts (K2-era ids and kimi-latest 404), matching the
|
||||
# management catalog.
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cached_input_per_1k: 0.0003
|
||||
|
||||
anthropic:
|
||||
# Claude 4.x family — cache reads ≈10% of input, cache writes ≈125% of input.
|
||||
# Pricing source: Anthropic's current published rates per million tokens,
|
||||
# divided by 1000 for the per-1k figures stored here.
|
||||
claude-fable-5:
|
||||
input_per_1k: 0.010
|
||||
output_per_1k: 0.050
|
||||
cache_read_per_1k: 0.001
|
||||
cache_creation_per_1k: 0.0125
|
||||
claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
|
||||
# Kimi / Moonshot AI (kimi_api) via the Anthropic-compatible endpoint
|
||||
# (/anthropic/v1/messages — the official Claude Code setup). Same rates
|
||||
# as the OpenAI-shape entry above. "kimi-k3[1m]" is the model id some
|
||||
# Claude Code guides set for the 1M-context alias; priced identically so
|
||||
# cost metering doesn't silently skip those requests.
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
"kimi-k3[1m]":
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
|
||||
bedrock:
|
||||
# AWS Bedrock model ids, normalised by the request parser (cross-region
|
||||
# inference-profile prefix + version/throughput suffix stripped), e.g.
|
||||
# eu.anthropic.claude-sonnet-4-5-20250929-v1:0 -> anthropic.claude-sonnet-4-5.
|
||||
# Anthropic-on-Bedrock keeps the additive cache buckets (read ≈0.1x input,
|
||||
# write ≈1.25x input); Nova / Llama report no cache, so cost is input+output.
|
||||
anthropic.claude-opus-5:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-8:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-7:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-6:
|
||||
input_per_1k: 0.005
|
||||
output_per_1k: 0.025
|
||||
cache_read_per_1k: 0.0005
|
||||
cache_creation_per_1k: 0.00625
|
||||
anthropic.claude-opus-4-1:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
cache_read_per_1k: 0.0015
|
||||
cache_creation_per_1k: 0.01875
|
||||
anthropic.claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-haiku-4-5:
|
||||
input_per_1k: 0.001
|
||||
output_per_1k: 0.005
|
||||
cache_read_per_1k: 0.0001
|
||||
cache_creation_per_1k: 0.00125
|
||||
meta.llama3-3-70b-instruct:
|
||||
input_per_1k: 0.00072
|
||||
output_per_1k: 0.00072
|
||||
amazon.nova-2-lite:
|
||||
input_per_1k: 0.0003
|
||||
output_per_1k: 0.0025
|
||||
amazon.nova-pro:
|
||||
input_per_1k: 0.0008
|
||||
output_per_1k: 0.0032
|
||||
amazon.nova-lite:
|
||||
mistral-large-latest:
|
||||
input_per_1k: 0.0005
|
||||
output_per_1k: 0.0015
|
||||
mistral-medium-3-5:
|
||||
input_per_1k: 0.0015
|
||||
output_per_1k: 0.0075
|
||||
mistral-medium-latest:
|
||||
input_per_1k: 0.0004
|
||||
output_per_1k: 0.002
|
||||
mistral-small-latest:
|
||||
input_per_1k: 0.00006
|
||||
output_per_1k: 0.00024
|
||||
amazon.nova-micro:
|
||||
input_per_1k: 0.000035
|
||||
output_per_1k: 0.00014
|
||||
output_per_1k: 0.00018
|
||||
o4-mini:
|
||||
input_per_1k: 0.0011
|
||||
output_per_1k: 0.0044
|
||||
cached_input_per_1k: 0.000275
|
||||
text-embedding-3-large:
|
||||
input_per_1k: 0.00013
|
||||
output_per_1k: 0
|
||||
text-embedding-3-small:
|
||||
input_per_1k: 0.00002
|
||||
output_per_1k: 0
|
||||
@@ -0,0 +1,148 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
)
|
||||
|
||||
// TestDefaultTable_CoversEveryCatalogModel replaces the proxy's old
|
||||
// hand-maintained coverage list: because the table is built FROM the
|
||||
// catalog, drift is impossible by construction — this test guards the
|
||||
// fold itself (every catalog model of every surfaced provider resolves,
|
||||
// with exactly the catalog's rates).
|
||||
func TestDefaultTable_CoversEveryCatalogModel(t *testing.T) {
|
||||
table := DefaultTable()
|
||||
for _, p := range catalog.All() {
|
||||
if len(p.PricingSurfaces) == 0 {
|
||||
assert.Empty(t, p.Models, "catalog entry %s declares models but no pricing surfaces — those models would never be priced", p.ID)
|
||||
continue
|
||||
}
|
||||
for _, surface := range p.PricingSurfaces {
|
||||
byModel, ok := table[surface]
|
||||
require.True(t, ok, "surface %q (provider %s) missing from default table", surface, p.ID)
|
||||
for _, m := range p.Models {
|
||||
e, ok := byModel[m.ID]
|
||||
require.True(t, ok, "%s/%s (provider %s) missing from default table", surface, m.ID, p.ID)
|
||||
assert.Equal(t, m.InputPer1k, e.InputPer1k, "%s/%s input rate", surface, m.ID)
|
||||
assert.Equal(t, m.OutputPer1k, e.OutputPer1k, "%s/%s output rate", surface, m.ID)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDefaultTable_NoConflictingContributions enforces the collision rule
|
||||
// documented on catalog.Provider.PricingSurfaces: when two catalog
|
||||
// providers contribute the same (surface, model) pair — azure/vertex
|
||||
// mirroring openai/anthropic, kimi on both surfaces — their rates must be
|
||||
// identical, because the surface-keyed table can only hold one entry.
|
||||
// If a provider ever diverges (e.g. Azure reprices a model), this fails
|
||||
// and the divergence must move to per-provider-record pricing.
|
||||
func TestDefaultTable_NoConflictingContributions(t *testing.T) {
|
||||
type contribution struct {
|
||||
providerID string
|
||||
entry Entry
|
||||
}
|
||||
seen := map[string]map[string]contribution{}
|
||||
for _, p := range catalog.All() {
|
||||
for _, surface := range p.PricingSurfaces {
|
||||
if seen[surface] == nil {
|
||||
seen[surface] = map[string]contribution{}
|
||||
}
|
||||
for _, m := range p.Models {
|
||||
e := entryFromCatalogModel(m)
|
||||
if prev, dup := seen[surface][m.ID]; dup {
|
||||
assert.Equal(t, prev.entry, e,
|
||||
"%s/%s: %s and %s contribute different rates", surface, m.ID, prev.providerID, p.ID)
|
||||
continue
|
||||
}
|
||||
seen[surface][m.ID] = contribution{providerID: p.ID, entry: e}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Supplemental entries must never shadow a catalog-contributed model —
|
||||
// they exist precisely because the catalog does NOT list them.
|
||||
for surface, models := range supplementalDefaults {
|
||||
for id := range models {
|
||||
_, fromCatalog := seen[surface][id]
|
||||
assert.False(t, fromCatalog, "supplemental %s/%s is now in the catalog — delete the supplemental row", surface, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDefaultTable_AllRatesFiniteNonNegative mirrors the proxy-side
|
||||
// NewTable validation so a bad catalog edit is caught here, at unit-test
|
||||
// time, rather than as a chain-build failure on every proxy.
|
||||
func TestDefaultTable_AllRatesFiniteNonNegative(t *testing.T) {
|
||||
for surface, models := range DefaultTable() {
|
||||
for id, e := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input": e.InputPer1k,
|
||||
"output": e.OutputPer1k,
|
||||
"cached_input": e.CachedInputPer1k,
|
||||
"cache_read": e.CacheReadPer1k,
|
||||
"cache_creation": e.CacheCreationPer1k,
|
||||
} {
|
||||
assert.False(t, v < 0 || math.IsNaN(v) || math.IsInf(v, 0),
|
||||
"%s/%s: %s rate %v must be finite and non-negative", surface, id, field, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDefaultTable_PinnedRates pins rates that previously drifted or are
|
||||
// easy to mis-enter (carried over from the proxy's retired
|
||||
// defaults_coverage_test), plus the supplemental entries.
|
||||
func TestDefaultTable_PinnedRates(t *testing.T) {
|
||||
table := DefaultTable()
|
||||
|
||||
gpt54 := table["openai"]["gpt-5.4"]
|
||||
assert.InDelta(t, 0.0025, gpt54.InputPer1k, 1e-9, "gpt-5.4 input")
|
||||
assert.InDelta(t, 0.015, gpt54.OutputPer1k, 1e-9, "gpt-5.4 output")
|
||||
assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "gpt-5.4 cached input")
|
||||
|
||||
sonnet := table["bedrock"]["anthropic.claude-sonnet-4-5"]
|
||||
assert.InDelta(t, 0.003, sonnet.InputPer1k, 1e-9, "bedrock sonnet-4-5 input")
|
||||
assert.InDelta(t, 0.015, sonnet.OutputPer1k, 1e-9, "bedrock sonnet-4-5 output")
|
||||
assert.InDelta(t, 0.0003, sonnet.CacheReadPer1k, 1e-9, "bedrock sonnet-4-5 cache read")
|
||||
assert.InDelta(t, 0.00375, sonnet.CacheCreationPer1k, 1e-9, "bedrock sonnet-4-5 cache creation")
|
||||
|
||||
// Vertex Claude prices under "anthropic" with the bare id.
|
||||
fable := table["anthropic"]["claude-fable-5"]
|
||||
assert.InDelta(t, 0.010, fable.InputPer1k, 1e-9, "claude-fable-5 input")
|
||||
assert.InDelta(t, 0.0125, fable.CacheCreationPer1k, 1e-9, "claude-fable-5 cache creation")
|
||||
|
||||
// Supplementals present on their surfaces.
|
||||
for surface, ids := range map[string][]string{
|
||||
"openai": {"gpt-5", "gpt-5-mini", "gpt-5-nano"},
|
||||
"anthropic": {"claude-opus-5", "kimi-k3[1m]", "kimi-k3"},
|
||||
"bedrock": {"anthropic.claude-opus-5"},
|
||||
} {
|
||||
for _, id := range ids {
|
||||
_, ok := table[surface][id]
|
||||
assert.True(t, ok, "%s/%s must be priced", surface, id)
|
||||
}
|
||||
}
|
||||
|
||||
// Embeddings bill input-only — output stays zero.
|
||||
emb := table["openai"]["text-embedding-3-large"]
|
||||
assert.Zero(t, emb.OutputPer1k, "embedding output rate must be zero")
|
||||
assert.Positive(t, emb.InputPer1k, "embedding input rate must be set")
|
||||
}
|
||||
|
||||
func TestLookupDefault_SurfaceOrder(t *testing.T) {
|
||||
// kimi-k3 exists on both surfaces; first surface in the slice wins.
|
||||
e, ok := LookupDefault([]string{"openai", "anthropic"}, "kimi-k3")
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.003, e.InputPer1k, 1e-9)
|
||||
|
||||
_, ok = LookupDefault([]string{"bedrock"}, "gpt-4o")
|
||||
assert.False(t, ok, "gpt-4o is not a bedrock model")
|
||||
|
||||
_, ok = LookupDefault(nil, "gpt-4o")
|
||||
assert.False(t, ok, "no surfaces, no match")
|
||||
}
|
||||
109
management/internals/modules/agentnetwork/pricing/exampleyaml.go
Normal file
109
management/internals/modules/agentnetwork/pricing/exampleyaml.go
Normal file
@@ -0,0 +1,109 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// MarshalDefaultsYAML renders the built-in default pricing table (catalog
|
||||
// + supplementals, WITHOUT any operator override) as the YAML schema
|
||||
// LoadOverrideFile consumes. It backs the generated
|
||||
// defaults_llm_pricing.example.yaml so operators start from a file that
|
||||
// matches the compiled-in rates exactly; a golden test keeps the two in
|
||||
// sync. Output is deterministic (sorted surfaces and models).
|
||||
func MarshalDefaultsYAML() []byte {
|
||||
var b bytes.Buffer
|
||||
b.WriteString(exampleHeader)
|
||||
|
||||
table := buildDefaultTable()
|
||||
surfaces := make([]string, 0, len(table))
|
||||
for s := range table {
|
||||
surfaces = append(surfaces, s)
|
||||
}
|
||||
sort.Strings(surfaces)
|
||||
|
||||
for _, surface := range surfaces {
|
||||
fmt.Fprintf(&b, "\n%s:\n", surface)
|
||||
models := table[surface]
|
||||
ids := make([]string, 0, len(models))
|
||||
for id := range models {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Strings(ids)
|
||||
for _, id := range ids {
|
||||
e := models[id]
|
||||
fmt.Fprintf(&b, " %s:\n", yamlKey(id))
|
||||
fmt.Fprintf(&b, " input_per_1k: %s\n", rate(e.InputPer1k))
|
||||
fmt.Fprintf(&b, " output_per_1k: %s\n", rate(e.OutputPer1k))
|
||||
if e.CachedInputPer1k > 0 {
|
||||
fmt.Fprintf(&b, " cached_input_per_1k: %s\n", rate(e.CachedInputPer1k))
|
||||
}
|
||||
if e.CacheReadPer1k > 0 {
|
||||
fmt.Fprintf(&b, " cache_read_per_1k: %s\n", rate(e.CacheReadPer1k))
|
||||
}
|
||||
if e.CacheCreationPer1k > 0 {
|
||||
fmt.Fprintf(&b, " cache_creation_per_1k: %s\n", rate(e.CacheCreationPer1k))
|
||||
}
|
||||
}
|
||||
}
|
||||
return b.Bytes()
|
||||
}
|
||||
|
||||
// rate renders a USD-per-1k rate without float noise ("0.00015", not
|
||||
// "0.000150000000...").
|
||||
func rate(v float64) string {
|
||||
return strconv.FormatFloat(v, 'f', -1, 64)
|
||||
}
|
||||
|
||||
// yamlKey quotes model ids that YAML would otherwise misparse (e.g.
|
||||
// "kimi-k3[1m]" starts a flow sequence unquoted).
|
||||
func yamlKey(id string) string {
|
||||
for _, r := range id {
|
||||
switch r {
|
||||
case '[', ']', '{', '}', ':', '#', ',', '&', '*', '!', '|', '>', '\'', '"', '%', '@', '`':
|
||||
return strconv.Quote(id)
|
||||
}
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
const exampleHeader = `# Default LLM pricing used by NetBird's Agent Network cost metering.
|
||||
# GENERATED from the management catalog — do not edit this file in the
|
||||
# repository; regenerate with:
|
||||
#
|
||||
# go generate ./management/internals/modules/agentnetwork/pricing
|
||||
#
|
||||
# Operators: copy this file to <datadir>/defaults_llm_pricing.yaml (or
|
||||
# any path configured via management.json:
|
||||
#
|
||||
# { "AgentNetwork": { "PricingDefaultsFile": "/path/defaults_llm_pricing.yaml" } }
|
||||
#
|
||||
# ) and adjust the entries you want to change. Management re-reads the
|
||||
# file periodically (mtime poll, every minute): the live table feeds the
|
||||
# proxies' cost metering and the dashboard's model-price prefill, so
|
||||
# edits apply without a restart. Your file only needs the entries you
|
||||
# want to change — but each entry REPLACES the built-in entry for that
|
||||
# surface+model whole, so repeat the cache rates you want to keep.
|
||||
# Unknown fields and negative or non-finite rates are rejected: at
|
||||
# startup that fails boot (for an explicitly configured path); at
|
||||
# runtime the previous table is kept and a warning is logged. Deleting
|
||||
# the file reverts to the built-in defaults below.
|
||||
#
|
||||
# Top-level keys are pricing surfaces — the parser shape requests are
|
||||
# metered under: "openai" (also Azure, Mistral, and OpenAI-compatible
|
||||
# gateways), "anthropic" (also Anthropic-on-Vertex), "bedrock"
|
||||
# (normalized ids, e.g. anthropic.claude-sonnet-4-5). Model keys must be
|
||||
# the normalized id the proxy meters (version/region suffixes stripped).
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Optional cache fields:
|
||||
# cached_input_per_1k OpenAI shape: rate for cached prompt tokens
|
||||
# (a SUBSET of input tokens). Absent -> cached
|
||||
# portion bills at input_per_1k.
|
||||
# cache_read_per_1k Anthropic shape: rate for cache_read tokens
|
||||
# (ADDITIVE to input). Absent -> input rate.
|
||||
# cache_creation_per_1k Anthropic shape: rate for cache_creation
|
||||
# tokens (ADDITIVE to input). Absent -> input
|
||||
# rate.
|
||||
`
|
||||
20
management/internals/modules/agentnetwork/pricing/gen.go
Normal file
20
management/internals/modules/agentnetwork/pricing/gen.go
Normal file
@@ -0,0 +1,20 @@
|
||||
//go:build ignore
|
||||
|
||||
// Regenerates defaults_llm_pricing.example.yaml from the compiled-in
|
||||
// default pricing table. Run via:
|
||||
//
|
||||
// go generate ./management/internals/modules/agentnetwork/pricing
|
||||
package main
|
||||
|
||||
import (
|
||||
"log"
|
||||
"os"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
)
|
||||
|
||||
func main() {
|
||||
if err := os.WriteFile("defaults_llm_pricing.example.yaml", pricing.MarshalDefaultsYAML(), 0o644); err != nil {
|
||||
log.Fatalf("write defaults_llm_pricing.example.yaml: %v", err)
|
||||
}
|
||||
}
|
||||
244
management/internals/modules/agentnetwork/pricing/override.go
Normal file
244
management/internals/modules/agentnetwork/pricing/override.go
Normal file
@@ -0,0 +1,244 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"math"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// DefaultFileName is the basename probed under management's datadir when
|
||||
// AgentNetwork.PricingDefaultsFile doesn't configure an explicit path.
|
||||
const DefaultFileName = "defaults_llm_pricing.yaml"
|
||||
|
||||
// ReloadInterval is the cadence at which the pricing file's mtime is
|
||||
// polled for changes.
|
||||
const ReloadInterval = time.Minute
|
||||
|
||||
// maxFileBytes bounds the pricing file read so a misconfigured path
|
||||
// (pointed at a huge file) cannot exhaust process memory.
|
||||
const maxFileBytes = 4 << 20
|
||||
|
||||
// pricingFile mirrors the on-disk YAML schema — the same schema the
|
||||
// proxy's retired embedded defaults_pricing.yaml used, so files written
|
||||
// for it keep working. Keys are pricing surfaces ("openai", "anthropic",
|
||||
// "bedrock"); nested keys are normalized model ids.
|
||||
type pricingFile map[string]map[string]struct {
|
||||
InputPer1k float64 `yaml:"input_per_1k"`
|
||||
OutputPer1k float64 `yaml:"output_per_1k"`
|
||||
CachedInputPer1k float64 `yaml:"cached_input_per_1k"`
|
||||
CacheReadPer1k float64 `yaml:"cache_read_per_1k"`
|
||||
CacheCreationPer1k float64 `yaml:"cache_creation_per_1k"`
|
||||
}
|
||||
|
||||
// fileState tracks the watched pricing file across reloads.
|
||||
var fileState struct {
|
||||
mu sync.Mutex
|
||||
path string
|
||||
mtime int64
|
||||
}
|
||||
|
||||
// LoadFile loads the management-side pricing defaults file and makes it
|
||||
// the live table (merged entry-whole over the compiled-in fallback; see
|
||||
// DefaultTable). The path stays registered for the periodic reloader, so
|
||||
// later edits — or the file (re)appearing after deletion — are picked up
|
||||
// without a restart.
|
||||
//
|
||||
// required governs the missing-file case: true for an explicitly
|
||||
// configured path (a typo must fail startup rather than silently bill
|
||||
// with built-ins the operator believes they replaced), false for the
|
||||
// conventional <datadir>/defaults_llm_pricing.yaml probe (absent file =
|
||||
// compiled-in defaults, still watched in case it appears). A file that
|
||||
// exists but is malformed is always an error at load time.
|
||||
func LoadFile(path string, required bool) error {
|
||||
if path == "" {
|
||||
return nil
|
||||
}
|
||||
fileState.mu.Lock()
|
||||
fileState.path = path
|
||||
fileState.mu.Unlock()
|
||||
|
||||
table, mtime, err := readFile(path)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) && !required {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
storeFileTable(table, mtime)
|
||||
return nil
|
||||
}
|
||||
|
||||
// StartReloader launches the periodic mtime-poll goroutine for the file
|
||||
// registered by LoadFile. Runtime failures are lenient — a parse error
|
||||
// keeps the previously loaded table and logs a warning; a deleted file
|
||||
// reverts to the compiled-in defaults — so a mid-edit save can never
|
||||
// take pricing down. Returns immediately when no path was registered.
|
||||
func StartReloader(ctx context.Context, interval time.Duration) {
|
||||
fileState.mu.Lock()
|
||||
path := fileState.path
|
||||
fileState.mu.Unlock()
|
||||
if path == "" {
|
||||
return
|
||||
}
|
||||
if interval <= 0 {
|
||||
interval = ReloadInterval
|
||||
}
|
||||
go func() {
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
reload()
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// reload performs one mtime check + reload cycle.
|
||||
func reload() {
|
||||
fileState.mu.Lock()
|
||||
path, lastMtime := fileState.path, fileState.mtime
|
||||
fileState.mu.Unlock()
|
||||
|
||||
st, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// File removed (or not yet created): serve compiled-in
|
||||
// defaults and reset mtime so a future (re)appearance loads.
|
||||
if mergedTable.Swap(nil) != nil {
|
||||
log.Warnf("agent-network pricing defaults file %s removed; reverting to built-in defaults", path)
|
||||
}
|
||||
setMtime(0)
|
||||
return
|
||||
}
|
||||
log.Warnf("agent-network pricing defaults reload: stat %s: %v", path, err)
|
||||
return
|
||||
}
|
||||
if st.ModTime().UnixNano() == lastMtime {
|
||||
return
|
||||
}
|
||||
|
||||
table, mtime, err := readFile(path)
|
||||
if err != nil {
|
||||
// Keep the previously loaded table — never blank prices because
|
||||
// an operator saved mid-edit.
|
||||
log.Warnf("agent-network pricing defaults reload failed for %s (keeping previous table): %v", path, err)
|
||||
return
|
||||
}
|
||||
storeFileTable(table, mtime)
|
||||
log.Infof("agent-network pricing defaults reloaded from %s", path)
|
||||
}
|
||||
|
||||
func readFile(path string) (map[string]map[string]Entry, int64, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("open pricing defaults %s: %w", path, err)
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
|
||||
st, err := f.Stat()
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("stat pricing defaults %s: %w", path, err)
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(f, maxFileBytes+1))
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("read pricing defaults %s: %w", path, err)
|
||||
}
|
||||
if len(data) > maxFileBytes {
|
||||
return nil, 0, fmt.Errorf("pricing defaults %s exceeds %d bytes", path, maxFileBytes)
|
||||
}
|
||||
table, err := parsePricingYAML(data)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("parse pricing defaults %s: %w", path, err)
|
||||
}
|
||||
return table, st.ModTime().UnixNano(), nil
|
||||
}
|
||||
|
||||
// storeFileTable merges the parsed file over the compiled-in base and
|
||||
// publishes the result as the live table. File entries replace the
|
||||
// built-in entry for the same (surface, model) whole — they are not
|
||||
// field-merged — and surfaces/models the file doesn't mention keep the
|
||||
// built-in rates, so a partial file only needs the entries it changes.
|
||||
func storeFileTable(table map[string]map[string]Entry, mtime int64) {
|
||||
base := compiledBase()
|
||||
merged := make(map[string]map[string]Entry, len(base)+len(table))
|
||||
for surface, models := range base {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for id, e := range models {
|
||||
inner[id] = e
|
||||
}
|
||||
merged[surface] = inner
|
||||
}
|
||||
for surface, models := range table {
|
||||
inner, ok := merged[surface]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry, len(models))
|
||||
merged[surface] = inner
|
||||
}
|
||||
for id, e := range models {
|
||||
inner[id] = e
|
||||
}
|
||||
}
|
||||
mergedTable.Store(&merged)
|
||||
setMtime(mtime)
|
||||
}
|
||||
|
||||
func setMtime(v int64) {
|
||||
fileState.mu.Lock()
|
||||
fileState.mtime = v
|
||||
fileState.mu.Unlock()
|
||||
}
|
||||
|
||||
// parsePricingYAML decodes and validates the pricing YAML. Unknown
|
||||
// fields are rejected (typos surface instead of silently pricing at 0)
|
||||
// and every rate must be a finite, non-negative USD amount — the same
|
||||
// constraints the HTTP API enforces on operator per-provider prices.
|
||||
func parsePricingYAML(data []byte) (map[string]map[string]Entry, error) {
|
||||
dec := yaml.NewDecoder(bytes.NewReader(data))
|
||||
dec.KnownFields(true)
|
||||
|
||||
var raw pricingFile
|
||||
if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) {
|
||||
return nil, fmt.Errorf("decode yaml: %w", err)
|
||||
}
|
||||
|
||||
out := make(map[string]map[string]Entry, len(raw))
|
||||
for surface, models := range raw {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, e := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input_per_1k": e.InputPer1k,
|
||||
"output_per_1k": e.OutputPer1k,
|
||||
"cached_input_per_1k": e.CachedInputPer1k,
|
||||
"cache_read_per_1k": e.CacheReadPer1k,
|
||||
"cache_creation_per_1k": e.CacheCreationPer1k,
|
||||
} {
|
||||
if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return nil, fmt.Errorf("%s/%s: %s must be a finite, non-negative rate, got %v", surface, model, field, v)
|
||||
}
|
||||
}
|
||||
inner[model] = Entry{
|
||||
InputPer1k: e.InputPer1k,
|
||||
OutputPer1k: e.OutputPer1k,
|
||||
CachedInputPer1k: e.CachedInputPer1k,
|
||||
CacheReadPer1k: e.CacheReadPer1k,
|
||||
CacheCreationPer1k: e.CacheCreationPer1k,
|
||||
}
|
||||
}
|
||||
out[surface] = inner
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// resetFileState snapshots and restores the package-level file state so
|
||||
// tests stay order-independent.
|
||||
func resetFileState(t *testing.T) {
|
||||
t.Helper()
|
||||
prevMerged := mergedTable.Load()
|
||||
fileState.mu.Lock()
|
||||
prevPath, prevMtime := fileState.path, fileState.mtime
|
||||
fileState.mu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
mergedTable.Store(prevMerged)
|
||||
fileState.mu.Lock()
|
||||
fileState.path, fileState.mtime = prevPath, prevMtime
|
||||
fileState.mu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
func writePricing(t *testing.T, path, yml string) {
|
||||
t.Helper()
|
||||
require.NoError(t, os.WriteFile(path, []byte(yml), 0o600))
|
||||
}
|
||||
|
||||
func TestLoadFile_MergesOverCompiledDefaults(t *testing.T) {
|
||||
resetFileState(t)
|
||||
path := filepath.Join(t.TempDir(), DefaultFileName)
|
||||
writePricing(t, path, `
|
||||
openai:
|
||||
# Reprice a built-in model. The entry replaces the built-in WHOLE:
|
||||
# omitting the cache rate here drops the built-in 0.00125 discount.
|
||||
gpt-4o:
|
||||
input_per_1k: 0.9
|
||||
output_per_1k: 1.8
|
||||
# A model NetBird doesn't know at all.
|
||||
my-private-ft:
|
||||
input_per_1k: 0.01
|
||||
output_per_1k: 0.02
|
||||
cached_input_per_1k: 0.005
|
||||
gemini:
|
||||
gemini-pro:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.005
|
||||
`)
|
||||
require.NoError(t, LoadFile(path, true))
|
||||
table := DefaultTable()
|
||||
|
||||
gpt4o := table["openai"]["gpt-4o"]
|
||||
assert.InDelta(t, 0.9, gpt4o.InputPer1k, 1e-9, "file rate replaces the compiled-in rate")
|
||||
assert.Zero(t, gpt4o.CachedInputPer1k, "entries replace whole — omitted cache rate is dropped, not inherited")
|
||||
|
||||
ft := table["openai"]["my-private-ft"]
|
||||
assert.InDelta(t, 0.005, ft.CachedInputPer1k, 1e-9, "unknown models are added to the surface")
|
||||
_, ok := table["gemini"]["gemini-pro"]
|
||||
assert.True(t, ok, "a surface the catalog doesn't declare can be added")
|
||||
|
||||
// Untouched entries keep compiled-in rates (catalog, other surface,
|
||||
// supplemental).
|
||||
assert.InDelta(t, 0.00015, table["openai"]["gpt-4o-mini"].InputPer1k, 1e-9, "unlisted model keeps compiled rate")
|
||||
assert.InDelta(t, 0.003, table["anthropic"]["claude-sonnet-4-5"].InputPer1k, 1e-9, "unlisted surface untouched")
|
||||
assert.InDelta(t, 0.00125, table["openai"]["gpt-5"].InputPer1k, 1e-9, "supplemental entries untouched")
|
||||
|
||||
// The synthesizer-facing lookup reads the live table too.
|
||||
e, ok := LookupDefault([]string{"openai"}, "gpt-4o")
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.9, e.InputPer1k, 1e-9, "LookupDefault serves the file-backed rate")
|
||||
}
|
||||
|
||||
func TestLoadFile_MissingPath(t *testing.T) {
|
||||
resetFileState(t)
|
||||
missing := filepath.Join(t.TempDir(), DefaultFileName)
|
||||
|
||||
require.Error(t, LoadFile(missing, true),
|
||||
"explicitly configured path that doesn't exist must fail startup")
|
||||
|
||||
require.NoError(t, LoadFile(missing, false),
|
||||
"conventional datadir probe tolerates an absent file (compiled-in defaults serve)")
|
||||
assert.Nil(t, mergedTable.Load(), "no file, no merged table")
|
||||
fileState.mu.Lock()
|
||||
path := fileState.path
|
||||
fileState.mu.Unlock()
|
||||
assert.Equal(t, missing, path, "the path stays registered so the reloader picks the file up when it appears")
|
||||
}
|
||||
|
||||
func TestLoadFile_RejectsInvalid(t *testing.T) {
|
||||
resetFileState(t)
|
||||
dir := t.TempDir()
|
||||
cases := map[string]string{
|
||||
"unknown field (typo)": "openai:\n gpt-4o:\n input_per1k: 0.1\n",
|
||||
"negative rate": "openai:\n gpt-4o:\n input_per_1k: -0.1\n",
|
||||
"non-numeric rate": "openai:\n gpt-4o:\n input_per_1k: cheap\n",
|
||||
"not a mapping": "- just\n- a\n- list\n",
|
||||
}
|
||||
for name, yml := range cases {
|
||||
path := filepath.Join(dir, DefaultFileName)
|
||||
writePricing(t, path, yml)
|
||||
assert.Error(t, LoadFile(path, true), "case %q must be rejected", name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReload_LifeCycle drives the reloader's single-shot reload through
|
||||
// its full lifecycle: file edit picked up on mtime change, a broken save
|
||||
// keeps the previous table, and file removal reverts to the compiled-in
|
||||
// defaults (then a re-created file loads again).
|
||||
func TestReload_LifeCycle(t *testing.T) {
|
||||
resetFileState(t)
|
||||
path := filepath.Join(t.TempDir(), DefaultFileName)
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.5\n output_per_1k: 1\n")
|
||||
require.NoError(t, LoadFile(path, true))
|
||||
require.InDelta(t, 0.5, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9)
|
||||
|
||||
// Edit: new mtime, new rates.
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.7\n output_per_1k: 1.4\n")
|
||||
bumpMtime(t, path)
|
||||
reload()
|
||||
assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9, "edit must be picked up")
|
||||
|
||||
// Broken save: previous table survives.
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: -1\n")
|
||||
bumpMtime(t, path)
|
||||
reload()
|
||||
assert.InDelta(t, 0.7, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9,
|
||||
"a malformed save must keep the previously loaded table, never blank prices")
|
||||
|
||||
// Removal: compiled-in defaults serve again.
|
||||
require.NoError(t, os.Remove(path))
|
||||
reload()
|
||||
assert.InDelta(t, 0.0025, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9,
|
||||
"file removal reverts to the compiled-in rate")
|
||||
|
||||
// Re-created file loads without a restart.
|
||||
writePricing(t, path, "openai:\n gpt-4o:\n input_per_1k: 0.9\n output_per_1k: 1.8\n")
|
||||
bumpMtime(t, path)
|
||||
reload()
|
||||
assert.InDelta(t, 0.9, DefaultTable()["openai"]["gpt-4o"].InputPer1k, 1e-9,
|
||||
"a file appearing after removal (or after a missing-probe boot) must load")
|
||||
}
|
||||
|
||||
// bumpMtime pushes the file's mtime forward past the previously recorded
|
||||
// value — timestamps can otherwise collide within the test's timescale.
|
||||
func bumpMtime(t *testing.T, path string) {
|
||||
t.Helper()
|
||||
st, err := os.Stat(path)
|
||||
require.NoError(t, err)
|
||||
next := st.ModTime().Add(2 * 1e9)
|
||||
require.NoError(t, os.Chtimes(path, next, next))
|
||||
}
|
||||
|
||||
// TestExampleYAML_InSyncWithBuiltins is the golden guard for
|
||||
// defaults_llm_pricing.example.yaml: the shipped example must stay
|
||||
// byte-identical to what the compiled-in table renders (catalog edits
|
||||
// require `go generate ./management/internals/modules/agentnetwork/pricing`)
|
||||
// and must round-trip through the same parser operators' files go
|
||||
// through, reproducing the compiled-in table exactly.
|
||||
func TestExampleYAML_InSyncWithBuiltins(t *testing.T) {
|
||||
onDisk, err := os.ReadFile("defaults_llm_pricing.example.yaml")
|
||||
require.NoError(t, err, "example file must exist next to the package")
|
||||
require.Equal(t, string(MarshalDefaultsYAML()), string(onDisk),
|
||||
"defaults_llm_pricing.example.yaml is stale — run: go generate ./management/internals/modules/agentnetwork/pricing")
|
||||
|
||||
parsed, err := parsePricingYAML(onDisk)
|
||||
require.NoError(t, err, "the example must be a valid pricing defaults file")
|
||||
assert.Equal(t, buildDefaultTable(), parsed,
|
||||
"parsing the example must reproduce the compiled-in table exactly")
|
||||
}
|
||||
@@ -233,9 +233,19 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
return nil, err
|
||||
}
|
||||
|
||||
costMeterJSON, err := buildCostMeterConfigJSON(enabledProviders, groupIndex)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
|
||||
applyAccountCollectionControls(&mergedGuardrails, settings)
|
||||
guardrailJSON, err := marshalGuardrailConfig(mergedGuardrails)
|
||||
// The proxy guardrail is a per-provider fail-closed backstop; the
|
||||
// authoritative per-policy/group decision is management's
|
||||
// SelectPolicyForRequest. A provider lands in this map only when every
|
||||
// authorising policy restricts models.
|
||||
providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID)
|
||||
guardrailJSON, err := marshalGuardrailConfig(providerAllowlists, mergedGuardrails.PromptCapture)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -243,7 +253,7 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
// Use the merged decision (account settings OR policy-required redaction),
|
||||
// not the raw account flag, so a policy that mandates PII redaction is
|
||||
// honored by the capture parsers even when the account toggle is off.
|
||||
middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled)
|
||||
middlewares := buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON, mergedGuardrails.PromptCapture.RedactPii, mergedGuardrails.PromptCapture.Enabled)
|
||||
|
||||
priv, pub, err := pickServiceSessionKeys(enabledProviders)
|
||||
if err != nil {
|
||||
@@ -695,7 +705,7 @@ func buildIdentityExtraHeaders(p *types.Provider, extras []catalog.ExtraHeader)
|
||||
// requests bound for gateways like LiteLLM that key budgets and
|
||||
// attribution off request headers. CanMutate is required so its
|
||||
// HeadersAdd / HeadersRemove pass the framework's mutation gate.
|
||||
func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig {
|
||||
func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON, costMeterJSON []byte, redactPii, capturePromptContent bool) []rpservice.MiddlewareConfig {
|
||||
// Both parsers receive an explicit capture flag derived from the account's
|
||||
// enable_prompt_collection toggle; nil/unset would default to the legacy
|
||||
// "always emit" behavior in the middleware, which is precisely what we
|
||||
@@ -764,10 +774,13 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt
|
||||
ConfigJSON: []byte("{}"),
|
||||
},
|
||||
{
|
||||
// Carries the full pricing table (defaults + per-provider
|
||||
// operator prices) so the proxy bills without an embedded
|
||||
// price list; see buildCostMeterConfigJSON.
|
||||
ID: middlewareIDCostMeter,
|
||||
Enabled: true,
|
||||
Slot: rpservice.MiddlewareSlotOnResponse,
|
||||
ConfigJSON: []byte("{}"),
|
||||
ConfigJSON: costMeterJSON,
|
||||
},
|
||||
{
|
||||
ID: middlewareIDLLMResponseParser,
|
||||
@@ -780,10 +793,12 @@ func buildMiddlewareChain(routerCfgJSON, identityInjectJSON, guardrailJSON []byt
|
||||
|
||||
// guardrailConfig is the JSON shape the proxy-side llm_guardrail
|
||||
// middleware expects. Mirrors the proxy registration documented in
|
||||
// the management→proxy contract.
|
||||
// the management→proxy contract. provider_allowlists is keyed by the
|
||||
// resolved provider id llm_router stamps; a provider absent from the map is
|
||||
// unrestricted at the proxy layer.
|
||||
type guardrailConfig struct {
|
||||
ModelAllowlist []string `json:"model_allowlist,omitempty"`
|
||||
PromptCapture guardrailPromptCapture `json:"prompt_capture"`
|
||||
ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"`
|
||||
PromptCapture guardrailPromptCapture `json:"prompt_capture"`
|
||||
}
|
||||
|
||||
type guardrailPromptCapture struct {
|
||||
@@ -828,13 +843,10 @@ func applyAccountCollectionControls(merged *MergedGuardrails, settings *types.Se
|
||||
merged.PromptCapture.RedactPii = settings.RedactPii || merged.PromptCapture.RedactPii
|
||||
}
|
||||
|
||||
func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) {
|
||||
func marshalGuardrailConfig(providerAllowlists map[string][]string, capture MergedPromptCapture) ([]byte, error) {
|
||||
cfg := guardrailConfig{
|
||||
ModelAllowlist: merged.ModelAllowlist,
|
||||
PromptCapture: guardrailPromptCapture{
|
||||
Enabled: merged.PromptCapture.Enabled,
|
||||
RedactPii: merged.PromptCapture.RedactPii,
|
||||
},
|
||||
ProviderAllowlists: providerAllowlists,
|
||||
PromptCapture: guardrailPromptCapture(capture),
|
||||
}
|
||||
out, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
@@ -843,6 +855,74 @@ func marshalGuardrailConfig(merged MergedGuardrails) ([]byte, error) {
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// buildProviderAllowlists returns the proxy's per-provider backstop: a provider
|
||||
// is included only when every authorising policy restricts models (their union);
|
||||
// if any leaves it unrestricted it is omitted, so management decides per group.
|
||||
func buildProviderAllowlists(policies []*types.Policy, byID map[string]*types.Guardrail) map[string][]string {
|
||||
type providerAcc struct {
|
||||
models map[string]struct{}
|
||||
anyUnrestricted bool
|
||||
}
|
||||
accs := make(map[string]*providerAcc)
|
||||
for _, p := range policies {
|
||||
if p == nil {
|
||||
continue
|
||||
}
|
||||
restricted, models := policyModelAllowlist(p, byID)
|
||||
for _, providerID := range p.DestinationProviderIDs {
|
||||
if providerID == "" {
|
||||
continue
|
||||
}
|
||||
acc, ok := accs[providerID]
|
||||
if !ok {
|
||||
acc = &providerAcc{models: make(map[string]struct{})}
|
||||
accs[providerID] = acc
|
||||
}
|
||||
if !restricted {
|
||||
acc.anyUnrestricted = true
|
||||
continue
|
||||
}
|
||||
for _, m := range models {
|
||||
acc.models[m] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
out := make(map[string][]string, len(accs))
|
||||
for providerID, acc := range accs {
|
||||
if acc.anyUnrestricted {
|
||||
continue
|
||||
}
|
||||
models := make([]string, 0, len(acc.models))
|
||||
for m := range acc.models {
|
||||
models = append(models, m)
|
||||
}
|
||||
sort.Strings(models)
|
||||
out[providerID] = models
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// policyModelAllowlist reports whether a policy restricts models (has an
|
||||
// allowlist-enabled guardrail) and the union of allowed models. Models are
|
||||
// verbatim; the proxy factory lowercases/trims them at decode time.
|
||||
func policyModelAllowlist(p *types.Policy, byID map[string]*types.Guardrail) (bool, []string) {
|
||||
restricted := false
|
||||
var models []string
|
||||
for _, gID := range p.GuardrailIDs {
|
||||
g, ok := byID[gID]
|
||||
if !ok || g == nil || !g.Checks.ModelAllowlist.Enabled {
|
||||
continue
|
||||
}
|
||||
restricted = true
|
||||
for _, m := range g.Checks.ModelAllowlist.Models {
|
||||
if m != "" {
|
||||
models = append(models, m)
|
||||
}
|
||||
}
|
||||
}
|
||||
return restricted, models
|
||||
}
|
||||
|
||||
// buildAccountService composes the per-account gateway Service. The
|
||||
// target carries the noop placeholder URL — the router middleware
|
||||
// rewrites every request to the matched provider's upstream before the
|
||||
@@ -986,38 +1066,11 @@ func unionSourceGroups(policies []*types.Policy) []string {
|
||||
return out
|
||||
}
|
||||
|
||||
// MergedGuardrails is the JSON shape passed to the proxy via the
|
||||
// guardrail middleware's config_json. Mirrors the proxy-side
|
||||
// expectations and is intentionally distinct from
|
||||
// types.GuardrailChecks so we can evolve either side independently.
|
||||
// MergedGuardrails is the synthesiser's fold target. Only prompt capture is
|
||||
// merged here — the model allowlist is emitted per-provider, and
|
||||
// token/budget/retention moved onto Policy.Limits and account Settings.
|
||||
type MergedGuardrails struct {
|
||||
ModelAllowlist []string `json:"model_allowlist,omitempty"`
|
||||
TokenLimits MergedTokenLimits `json:"token_limits"`
|
||||
Budget MergedBudget `json:"budget"`
|
||||
PromptCapture MergedPromptCapture `json:"prompt_capture"`
|
||||
Retention MergedRetention `json:"retention"`
|
||||
}
|
||||
|
||||
type MergedTokenLimits struct {
|
||||
Hourly *MergedTokenWindow `json:"hourly,omitempty"`
|
||||
Daily *MergedTokenWindow `json:"daily,omitempty"`
|
||||
Monthly *MergedTokenWindow `json:"monthly,omitempty"`
|
||||
}
|
||||
|
||||
type MergedTokenWindow struct {
|
||||
MaxInputTokens int `json:"max_input_tokens,omitempty"`
|
||||
MaxOutputTokens int `json:"max_output_tokens,omitempty"`
|
||||
}
|
||||
|
||||
type MergedBudget struct {
|
||||
Hourly *MergedBudgetWindow `json:"hourly,omitempty"`
|
||||
Daily *MergedBudgetWindow `json:"daily,omitempty"`
|
||||
Monthly *MergedBudgetWindow `json:"monthly,omitempty"`
|
||||
}
|
||||
|
||||
type MergedBudgetWindow struct {
|
||||
SoftCapUSD float64 `json:"soft_cap_usd,omitempty"`
|
||||
HardCapUSD float64 `json:"hard_cap_usd,omitempty"`
|
||||
PromptCapture MergedPromptCapture
|
||||
}
|
||||
|
||||
type MergedPromptCapture struct {
|
||||
@@ -1025,64 +1078,31 @@ type MergedPromptCapture struct {
|
||||
RedactPii bool `json:"redact_pii"`
|
||||
}
|
||||
|
||||
type MergedRetention struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Days int `json:"days"`
|
||||
}
|
||||
|
||||
// mergeGuardrails computes the effective guardrail spec applied at the
|
||||
// proxy, given the referencing policies and the account's guardrail
|
||||
// catalogue. Policy enabled-ness is the caller's responsibility — only
|
||||
// enabled policies should be passed in.
|
||||
// mergeGuardrails folds the referencing policies' guardrails into the
|
||||
// prompt-capture decision only. The model allowlist is enforced per-policy/group
|
||||
// in management and shipped per-provider; token/budget/retention live off
|
||||
// guardrails now.
|
||||
//
|
||||
// Merge rules:
|
||||
// - Model allowlist: union of allowlists across policies that enable it.
|
||||
// - Token / Budget: most-restrictive (min of non-zero caps) per window.
|
||||
// - Prompt capture: enabled if any policy enables it; redact_pii sticks
|
||||
// if any enabling policy turns it on.
|
||||
// - Retention: enabled if any enables it; smallest non-zero days wins.
|
||||
// Merge rule — prompt capture: enabled if any policy enables it; redact_pii
|
||||
// sticks if any enabling policy turns it on.
|
||||
func mergeGuardrails(policies []*types.Policy, byID map[string]*types.Guardrail) MergedGuardrails {
|
||||
merged := MergedGuardrails{}
|
||||
allowlist := make(map[string]struct{})
|
||||
allowlistEnabled := false
|
||||
|
||||
for _, policy := range policies {
|
||||
for _, gID := range policy.GuardrailIDs {
|
||||
g, ok := byID[gID]
|
||||
if !ok || g == nil {
|
||||
continue
|
||||
}
|
||||
mergeGuardrail(g, &merged, allowlist, &allowlistEnabled)
|
||||
mergeGuardrail(g, &merged)
|
||||
}
|
||||
}
|
||||
|
||||
if allowlistEnabled {
|
||||
merged.ModelAllowlist = make([]string, 0, len(allowlist))
|
||||
for m := range allowlist {
|
||||
merged.ModelAllowlist = append(merged.ModelAllowlist, m)
|
||||
}
|
||||
sort.Strings(merged.ModelAllowlist)
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
// mergeGuardrail folds a single guardrail's enabled checks into the
|
||||
// running merge: model-allowlist models join the shared set (and flip
|
||||
// allowlistEnabled), and prompt-capture / redact-pii stick once any
|
||||
// enabling guardrail turns them on.
|
||||
//
|
||||
// TokenLimits, Budget, and Retention have moved off guardrails — token
|
||||
// and budget caps now live on the Policy itself (Policy.Limits) and
|
||||
// retention moves to account-level Settings — so they are not merged here.
|
||||
func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails, allowlist map[string]struct{}, allowlistEnabled *bool) {
|
||||
if g.Checks.ModelAllowlist.Enabled {
|
||||
*allowlistEnabled = true
|
||||
for _, m := range g.Checks.ModelAllowlist.Models {
|
||||
if m != "" {
|
||||
allowlist[m] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
// mergeGuardrail folds a single guardrail's prompt-capture settings into the
|
||||
// running merge: enabled / redact-pii stick once any enabling guardrail turns
|
||||
// them on.
|
||||
func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) {
|
||||
if g.Checks.PromptCapture.Enabled {
|
||||
merged.PromptCapture.Enabled = true
|
||||
if g.Checks.PromptCapture.RedactPii {
|
||||
|
||||
@@ -81,8 +81,8 @@ func TestSynthesizeServices_RealStore_PromptCaptureAccountIsSoleControl(t *testi
|
||||
require.Len(t, services, 1)
|
||||
|
||||
cfg := decodeServiceGuardrailConfig(t, services[0])
|
||||
assert.Equal(t, []string{"gpt-5.4"}, cfg.ModelAllowlist,
|
||||
"model allowlist is a pure policy guardrail and must always reach the config")
|
||||
assert.Equal(t, map[string][]string{"prov-1": {"gpt-5.4"}}, cfg.ProviderAllowlists,
|
||||
"model allowlist is a pure policy guardrail and must reach the per-provider config")
|
||||
assert.False(t, cfg.PromptCapture.Enabled,
|
||||
"prompt capture must be off when the account toggle is off, even with a capture-enabled guardrail")
|
||||
}
|
||||
@@ -172,7 +172,7 @@ func TestSynthesizeServices_RealStore_NoGuardrail_CaptureOff(t *testing.T) {
|
||||
require.Len(t, services, 1, "exactly one synth service expected")
|
||||
|
||||
cfg := decodeServiceGuardrailConfig(t, services[0])
|
||||
assert.Empty(t, cfg.ModelAllowlist, "no guardrail → no allowlist")
|
||||
assert.Empty(t, cfg.ProviderAllowlists, "no guardrail → provider unrestricted (absent from map)")
|
||||
assert.False(t, cfg.PromptCapture.Enabled, "no guardrail → prompt capture off by default")
|
||||
assert.False(t, cfg.PromptCapture.RedactPii, "no guardrail → redact off by default")
|
||||
}
|
||||
|
||||
131
management/internals/modules/agentnetwork/synthesizer_pricing.go
Normal file
131
management/internals/modules/agentnetwork/synthesizer_pricing.go
Normal file
@@ -0,0 +1,131 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
sharedllm "github.com/netbirdio/netbird/shared/llm"
|
||||
)
|
||||
|
||||
// costMeterConfig is the JSON shape the proxy-side cost_meter middleware
|
||||
// expects (mirror-type pattern, same as routerConfig). The top-level
|
||||
// "pricing" wrapper is the feature-detection signal: an old proxy's config
|
||||
// struct ignores it as an unknown field, and a new proxy treats its
|
||||
// absence as "old management" (skips every cost computation and warns).
|
||||
type costMeterConfig struct {
|
||||
Pricing *costMeterPricing `json:"pricing,omitempty"`
|
||||
}
|
||||
|
||||
// costMeterPricing carries the full pricing table:
|
||||
// - Defaults: surface ("openai"/"anthropic"/"bedrock") -> normalized
|
||||
// model id -> rates. The full default table ships to every account —
|
||||
// it is small (~10 KB) and keeps gateway-style providers (which
|
||||
// enumerate no models) priced for every catalog model.
|
||||
// - Providers: provider record id (matched against the
|
||||
// llm.resolved_provider_id metadata llm_router stamps) -> normalized
|
||||
// model id -> rates. Entries are fully materialized here at synth
|
||||
// time — default cache rates already folded in — so the proxy does
|
||||
// two map lookups and no merging.
|
||||
type costMeterPricing struct {
|
||||
Defaults map[string]map[string]pricing.Entry `json:"defaults,omitempty"`
|
||||
Providers map[string]map[string]pricing.Entry `json:"providers,omitempty"`
|
||||
}
|
||||
|
||||
// buildCostMeterConfigJSON assembles the cost_meter middleware config
|
||||
// from the default pricing table plus the operator's stored per-provider
|
||||
// model prices. Same orphan rule as the router: a provider no enabled
|
||||
// policy authorises is unreachable, so its prices are not shipped.
|
||||
//
|
||||
// Overlay semantics per model row:
|
||||
// - The row's model id is normalized exactly the way the proxy's
|
||||
// request parser normalizes the ids it meters (bedrock ARN/region/
|
||||
// version stripping, vertex "@version" stripping), so the per-record
|
||||
// lookup key compares equal to llm.model at billing time.
|
||||
// - The entry starts from the default entry for that model (when one
|
||||
// exists) to inherit cache rates the operator didn't state.
|
||||
// - Operator input/output overlay verbatim — including an explicit 0,
|
||||
// which prices the model as free (self-hosted / internal endpoints)
|
||||
// rather than silently reverting to list price.
|
||||
// - Cache-rate pointers overlay only when non-nil: nil means "inherit
|
||||
// the default", an explicit 0 means "no discount, bill this bucket
|
||||
// at the input rate".
|
||||
func buildCostMeterConfigJSON(providers []*types.Provider, groupIndex map[string][]string) ([]byte, error) {
|
||||
cfg := costMeterConfig{Pricing: &costMeterPricing{
|
||||
Defaults: pricing.DefaultTable(),
|
||||
}}
|
||||
|
||||
perRecord := make(map[string]map[string]pricing.Entry)
|
||||
for _, p := range providers {
|
||||
if _, hasPolicy := groupIndex[p.ID]; !hasPolicy {
|
||||
// Orphan: unreachable via the router, so unpriceable.
|
||||
continue
|
||||
}
|
||||
if len(p.Models) == 0 {
|
||||
// Gateway-style "claim every model" provider: the defaults
|
||||
// table is its price list.
|
||||
continue
|
||||
}
|
||||
entry, _ := catalog.Lookup(p.ProviderID)
|
||||
models := make(map[string]pricing.Entry, len(p.Models))
|
||||
for _, m := range p.Models {
|
||||
id := normalizePricingModelID(p.ProviderID, m.ID)
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
if _, dup := models[id]; dup {
|
||||
// First occurrence wins on post-normalization duplicates,
|
||||
// matching providerModelIDs' dedup order for routing.
|
||||
continue
|
||||
}
|
||||
models[id] = materializeEntry(entry.PricingSurfaces, id, m)
|
||||
}
|
||||
if len(models) > 0 {
|
||||
perRecord[p.ID] = models
|
||||
}
|
||||
}
|
||||
if len(perRecord) > 0 {
|
||||
cfg.Pricing.Providers = perRecord
|
||||
}
|
||||
|
||||
out, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal cost_meter middleware config: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// normalizePricingModelID maps an operator-entered model id onto the
|
||||
// normalized id the proxy's request parser emits as llm.model — the key
|
||||
// the cost meter looks up at billing time.
|
||||
func normalizePricingModelID(catalogProviderID, modelID string) string {
|
||||
switch {
|
||||
case catalog.IsBedrockPathStyle(catalogProviderID):
|
||||
return sharedllm.NormalizeBedrockModel(modelID)
|
||||
case catalog.IsVertexPathStyle(catalogProviderID):
|
||||
return sharedllm.NormalizeVertexModel(modelID)
|
||||
default:
|
||||
return modelID
|
||||
}
|
||||
}
|
||||
|
||||
// materializeEntry folds the default entry for (surfaces, model) — when
|
||||
// one exists — under the operator's stored prices, producing the fully
|
||||
// materialized wire entry.
|
||||
func materializeEntry(surfaces []string, normalizedID string, m types.ProviderModel) pricing.Entry {
|
||||
e, _ := pricing.LookupDefault(surfaces, normalizedID) // zero Entry on miss
|
||||
e.InputPer1k = m.InputPer1k
|
||||
e.OutputPer1k = m.OutputPer1k
|
||||
if m.CachedInputPer1k != nil {
|
||||
e.CachedInputPer1k = *m.CachedInputPer1k
|
||||
}
|
||||
if m.CacheReadPer1k != nil {
|
||||
e.CacheReadPer1k = *m.CacheReadPer1k
|
||||
}
|
||||
if m.CacheCreationPer1k != nil {
|
||||
e.CacheCreationPer1k = *m.CacheCreationPer1k
|
||||
}
|
||||
return e
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
)
|
||||
|
||||
func fptr(v float64) *float64 { return &v }
|
||||
|
||||
func decodeCostMeterConfig(t *testing.T, raw []byte) costMeterConfig {
|
||||
t.Helper()
|
||||
var cfg costMeterConfig
|
||||
require.NoError(t, json.Unmarshal(raw, &cfg), "cost meter config must round-trip")
|
||||
require.NotNil(t, cfg.Pricing, "pricing wrapper must be present")
|
||||
return cfg
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_BedrockModelNormalization: the operator may
|
||||
// paste region-prefixed, versioned, or ARN-wrapped Bedrock ids; the
|
||||
// per-record map must be keyed by the normalized id the request parser
|
||||
// emits as llm.model, or the lookup never hits at billing time.
|
||||
func TestBuildCostMeterConfig_BedrockModelNormalization(t *testing.T) {
|
||||
bedrock := &types.Provider{
|
||||
ID: "prov-bedrock",
|
||||
ProviderID: "bedrock_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{
|
||||
{ID: "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 0.0033, OutputPer1k: 0.0165},
|
||||
// Post-normalization duplicate of the row above under a
|
||||
// different regional spelling — first occurrence wins.
|
||||
{ID: "us.anthropic.claude-sonnet-4-5-20250929-v1:0", InputPer1k: 9.9, OutputPer1k: 9.9},
|
||||
},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON([]*types.Provider{bedrock}, map[string][]string{"prov-bedrock": {"grp"}})
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
|
||||
models := cfg.Pricing.Providers["prov-bedrock"]
|
||||
require.Len(t, models, 1, "both spellings normalize to one model; first row wins")
|
||||
e, ok := models["anthropic.claude-sonnet-4-5"]
|
||||
require.True(t, ok, "key must be the normalized id the parser emits, not the operator's raw spelling")
|
||||
assert.InDelta(t, 0.0033, e.InputPer1k, 1e-9, "first row's rate wins the dedup")
|
||||
assert.InDelta(t, 0.0003, e.CacheReadPer1k, 1e-9, "cache read inherited from the bedrock default entry")
|
||||
assert.InDelta(t, 0.00375, e.CacheCreationPer1k, 1e-9, "cache creation inherited from the bedrock default entry")
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_CacheRateNilVsZero pins the pointer semantics:
|
||||
// nil inherits the default cache rate, explicit 0 clears it (that bucket
|
||||
// bills at the input rate on the proxy).
|
||||
func TestBuildCostMeterConfig_CacheRateNilVsZero(t *testing.T) {
|
||||
p := &types.Provider{
|
||||
ID: "prov-oai",
|
||||
ProviderID: "openai_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{
|
||||
{ID: "gpt-4o", InputPer1k: 0.002, OutputPer1k: 0.008}, // nil → inherit 0.00125
|
||||
{ID: "gpt-4o-mini", InputPer1k: 0.0001, OutputPer1k: 0.0005, CachedInputPer1k: fptr(0)}, // explicit 0 → no discount
|
||||
{ID: "my-custom-ft", InputPer1k: 0.01, OutputPer1k: 0.02, CachedInputPer1k: fptr(0.005)}, // unknown model, explicit rate
|
||||
},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON([]*types.Provider{p}, map[string][]string{"prov-oai": {"grp"}})
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
models := cfg.Pricing.Providers["prov-oai"]
|
||||
|
||||
assert.InDelta(t, 0.00125, models["gpt-4o"].CachedInputPer1k, 1e-9, "nil cache pointer inherits the default rate")
|
||||
assert.Zero(t, models["gpt-4o-mini"].CachedInputPer1k, "explicit 0 overrides the default (0.000075) — bucket bills at input rate")
|
||||
custom := models["my-custom-ft"]
|
||||
assert.InDelta(t, 0.005, custom.CachedInputPer1k, 1e-9, "unknown model keeps the operator's explicit cache rate")
|
||||
assert.Zero(t, custom.CacheReadPer1k, "no default to inherit for a model outside the catalog")
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_OrphanAndGatewayProviders: an orphan (no
|
||||
// authorising policy) is unreachable so its prices must not ship; a
|
||||
// gateway with no model rows relies on the defaults table and gets no
|
||||
// per-record entry.
|
||||
func TestBuildCostMeterConfig_OrphanAndGatewayProviders(t *testing.T) {
|
||||
orphan := &types.Provider{
|
||||
ID: "prov-orphan",
|
||||
ProviderID: "openai_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{{ID: "gpt-4o", InputPer1k: 1, OutputPer1k: 1}},
|
||||
}
|
||||
gateway := &types.Provider{
|
||||
ID: "prov-litellm",
|
||||
ProviderID: "litellm_proxy",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON(
|
||||
[]*types.Provider{orphan, gateway},
|
||||
map[string][]string{"prov-litellm": {"grp"}}, // orphan has no policy
|
||||
)
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
|
||||
assert.NotContains(t, cfg.Pricing.Providers, "prov-orphan", "orphan provider prices must not ship")
|
||||
assert.NotContains(t, cfg.Pricing.Providers, "prov-litellm", "empty-models gateway needs no per-record entry")
|
||||
assert.NotEmpty(t, cfg.Pricing.Defaults["openai"], "defaults still ship so the gateway's catalog-model traffic is priced")
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
)
|
||||
|
||||
// policyForProviders builds an enabled policy authorising the given providers
|
||||
// under the given guardrails (both optional). Groups are irrelevant to
|
||||
// buildProviderAllowlists, which keys purely on destination provider.
|
||||
func policyForProviders(id string, guardrailIDs []string, providerIDs ...string) *types.Policy {
|
||||
return &types.Policy{
|
||||
ID: id,
|
||||
Enabled: true,
|
||||
DestinationProviderIDs: providerIDs,
|
||||
GuardrailIDs: guardrailIDs,
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildProviderAllowlists(t *testing.T) {
|
||||
byID := map[string]*types.Guardrail{
|
||||
"g-4o": allowlistGuardrail("g-4o", "acc-1", "gpt-4o"),
|
||||
"g-opus": allowlistGuardrail("g-opus", "acc-1", "claude-opus-4"),
|
||||
"g-disabled": {ID: "g-disabled", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}}},
|
||||
}
|
||||
|
||||
t.Run("all authorising policies restrict yields per-provider union", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
|
||||
policyForProviders("p2", []string{"g-opus"}, "prov-x"),
|
||||
}
|
||||
got := buildProviderAllowlists(policies, byID)
|
||||
assert.Equal(t, map[string][]string{"prov-x": {"claude-opus-4", "gpt-4o"}}, got,
|
||||
"a provider every policy restricts carries the sorted union of their models")
|
||||
})
|
||||
|
||||
t.Run("any un-guardrailed policy leaves the provider unrestricted (omitted)", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
|
||||
policyForProviders("p2", nil, "prov-x"), // no guardrail
|
||||
}
|
||||
got := buildProviderAllowlists(policies, byID)
|
||||
assert.NotContains(t, got, "prov-x",
|
||||
"a provider reachable by an un-guardrailed policy must be omitted so the proxy treats it as unrestricted")
|
||||
})
|
||||
|
||||
t.Run("a disabled allowlist counts as unrestricted", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForProviders("p1", []string{"g-disabled"}, "prov-x"),
|
||||
}
|
||||
got := buildProviderAllowlists(policies, byID)
|
||||
assert.NotContains(t, got, "prov-x",
|
||||
"a policy whose only guardrail has a disabled allowlist is unrestricted")
|
||||
})
|
||||
|
||||
t.Run("providers are isolated from one another", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForProviders("p1", []string{"g-4o"}, "prov-x"),
|
||||
policyForProviders("p2", []string{"g-opus"}, "prov-y"),
|
||||
}
|
||||
got := buildProviderAllowlists(policies, byID)
|
||||
assert.Equal(t, []string{"gpt-4o"}, got["prov-x"], "prov-x keeps only its own model")
|
||||
assert.Equal(t, []string{"claude-opus-4"}, got["prov-y"], "prov-y keeps only its own model")
|
||||
})
|
||||
|
||||
t.Run("one policy authorising two providers restricts both", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForProviders("p1", []string{"g-4o"}, "prov-x", "prov-y"),
|
||||
}
|
||||
got := buildProviderAllowlists(policies, byID)
|
||||
assert.Equal(t, []string{"gpt-4o"}, got["prov-x"])
|
||||
assert.Equal(t, []string{"gpt-4o"}, got["prov-y"])
|
||||
})
|
||||
|
||||
t.Run("union across a single policy's guardrails", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForProviders("p1", []string{"g-4o", "g-opus"}, "prov-x"),
|
||||
}
|
||||
got := buildProviderAllowlists(policies, byID)
|
||||
assert.ElementsMatch(t, []string{"claude-opus-4", "gpt-4o"}, got["prov-x"],
|
||||
"a policy's own multiple allowlist guardrails union together")
|
||||
})
|
||||
|
||||
t.Run("an enabled allowlist with no models denies everything", func(t *testing.T) {
|
||||
empty := map[string]*types.Guardrail{"g-empty": allowlistGuardrail("g-empty", "acc-1")}
|
||||
got := buildProviderAllowlists([]*types.Policy{
|
||||
policyForProviders("p1", []string{"g-empty"}, "prov-x"),
|
||||
}, empty)
|
||||
assert.Equal(t, map[string][]string{"prov-x": {}}, got,
|
||||
"an enabled-but-empty allowlist is restricted with an empty set, not unrestricted")
|
||||
})
|
||||
}
|
||||
@@ -33,14 +33,17 @@ func newSynthTestSettings() *types.Settings {
|
||||
|
||||
func newSynthTestProvider() *types.Provider {
|
||||
return &types.Provider{
|
||||
ID: "prov-1",
|
||||
AccountID: testAccountID,
|
||||
ProviderID: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test-key",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.0025, OutputPer1k: 0.015}},
|
||||
ID: "prov-1",
|
||||
AccountID: testAccountID,
|
||||
ProviderID: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test-key",
|
||||
Enabled: true,
|
||||
// Prices deliberately differ from the catalog's gpt-5.4 rates
|
||||
// (0.0025/0.015) so pricing tests can prove the operator's
|
||||
// stored price overlays the catalog default.
|
||||
Models: []types.ProviderModel{{ID: "gpt-5.4", InputPer1k: 0.004, OutputPer1k: 0.02}},
|
||||
SessionPrivateKey: "test-priv-key",
|
||||
SessionPublicKey: "test-pub-key",
|
||||
CreatedAt: time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC),
|
||||
@@ -214,7 +217,27 @@ func TestSynthesizeServices_HappyPath(t *testing.T) {
|
||||
|
||||
assert.Equal(t, middlewareIDCostMeter, mws[6].ID, "seventh middleware is the cost meter")
|
||||
assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[6].Slot, "cost meter runs on_response")
|
||||
assert.Equal(t, []byte("{}"), mws[6].ConfigJSON, "cost meter carries an explicit empty config")
|
||||
|
||||
var costCfg costMeterConfig
|
||||
require.NoError(t, json.Unmarshal(mws[6].ConfigJSON, &costCfg), "cost meter config must unmarshal")
|
||||
require.NotNil(t, costCfg.Pricing, "cost meter config must carry the pricing table — its absence tells the proxy management predates config-delivered pricing")
|
||||
|
||||
gpt4o, ok := costCfg.Pricing.Defaults["openai"]["gpt-4o"]
|
||||
require.True(t, ok, "the full default table ships regardless of the account's providers")
|
||||
assert.InDelta(t, 0.0025, gpt4o.InputPer1k, 1e-9, "default gpt-4o input rate comes from the catalog")
|
||||
|
||||
openaiPrices, ok := costCfg.Pricing.Providers[openai.ID]
|
||||
require.True(t, ok, "operator-priced provider must have a per-record entry")
|
||||
gpt54, ok := openaiPrices["gpt-5.4"]
|
||||
require.True(t, ok, "operator's model row keys the per-record map")
|
||||
assert.InDelta(t, 0.004, gpt54.InputPer1k, 1e-9, "operator input price overlays the catalog default (0.0025)")
|
||||
assert.InDelta(t, 0.02, gpt54.OutputPer1k, 1e-9, "operator output price overlays the catalog default (0.015)")
|
||||
assert.InDelta(t, 0.00025, gpt54.CachedInputPer1k, 1e-9, "cache rate the operator didn't state is inherited from the default entry")
|
||||
|
||||
opus, ok := costCfg.Pricing.Providers[anthropic.ID]["claude-opus-4-7"]
|
||||
require.True(t, ok, "anthropic's model row keys its per-record map")
|
||||
assert.Zero(t, opus.InputPer1k, "operator-stored zero prices ship verbatim — an explicit $0 model bills as free, it does not revert to list price")
|
||||
assert.InDelta(t, 0.0005, opus.CacheReadPer1k, 1e-9, "cache rates still inherit from the default entry")
|
||||
|
||||
assert.Equal(t, middlewareIDLLMResponseParser, mws[7].ID, "eighth middleware is the response parser")
|
||||
assert.Equal(t, rpservice.MiddlewareSlotOnResponse, mws[7].Slot, "response parser runs on_response")
|
||||
@@ -1031,8 +1054,12 @@ func TestSynthesizeServices_GuardrailMerge_AllowlistUnion_LimitsRestrictive(t *t
|
||||
|
||||
var cfg guardrailConfig
|
||||
require.NoError(t, json.Unmarshal(guardrailJSON, &cfg), "guardrail config must unmarshal cleanly")
|
||||
assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ModelAllowlist,
|
||||
"model allowlist union must keep both models")
|
||||
// Both policies restrict the same provider, so the per-provider backstop
|
||||
// carries the union of their models — a coarse gate that management's
|
||||
// per-policy/group check narrows; it only blocks models outside the union
|
||||
// when management is down.
|
||||
assert.ElementsMatch(t, []string{"gpt-5.4-mini", "gpt-5.4-pro"}, cfg.ProviderAllowlists["prov-1"],
|
||||
"per-provider allowlist union must keep both models")
|
||||
}
|
||||
|
||||
func TestSynthesizeServices_BackfillsMissingSessionKeys(t *testing.T) {
|
||||
|
||||
@@ -14,10 +14,24 @@ import (
|
||||
// ProviderModel is one row in the provider's models list. The operator
|
||||
// pins the per-1k input/output price for cost tracking; ID is the
|
||||
// model identifier the upstream provider expects on the wire.
|
||||
//
|
||||
// The three cache rates are pointers because absence is meaningful: nil
|
||||
// means "inherit NetBird's default rate for this model" (folded in at
|
||||
// synthesis time), while an explicit 0 means "no discount — bill this
|
||||
// cache bucket at the input rate".
|
||||
type ProviderModel struct {
|
||||
ID string `json:"id"`
|
||||
InputPer1k float64 `json:"input_per_1k"`
|
||||
OutputPer1k float64 `json:"output_per_1k"`
|
||||
// CachedInputPer1k is the OpenAI-shape rate for cached prompt tokens
|
||||
// (a subset of input tokens).
|
||||
CachedInputPer1k *float64 `json:"cached_input_per_1k,omitempty"`
|
||||
// CacheReadPer1k is the Anthropic-shape rate for cache-read tokens
|
||||
// (additive to input tokens).
|
||||
CacheReadPer1k *float64 `json:"cache_read_per_1k,omitempty"`
|
||||
// CacheCreationPer1k is the Anthropic-shape rate for cache-creation
|
||||
// tokens (additive to input tokens).
|
||||
CacheCreationPer1k *float64 `json:"cache_creation_per_1k,omitempty"`
|
||||
}
|
||||
|
||||
// Provider is an Agent Network AI provider record persisted per account.
|
||||
@@ -128,9 +142,12 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) {
|
||||
if req.Models != nil {
|
||||
for _, m := range *req.Models {
|
||||
p.Models = append(p.Models, ProviderModel{
|
||||
ID: m.Id,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
ID: m.Id,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k),
|
||||
CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k),
|
||||
CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -164,9 +181,12 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
|
||||
models := make([]api.AgentNetworkProviderModel, 0, len(p.Models))
|
||||
for _, m := range p.Models {
|
||||
models = append(models, api.AgentNetworkProviderModel{
|
||||
Id: m.ID,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
Id: m.ID,
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
CachedInputPer1k: copyFloatPtr(m.CachedInputPer1k),
|
||||
CacheReadPer1k: copyFloatPtr(m.CacheReadPer1k),
|
||||
CacheCreationPer1k: copyFloatPtr(m.CacheCreationPer1k),
|
||||
})
|
||||
}
|
||||
created := p.CreatedAt
|
||||
@@ -201,11 +221,27 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
|
||||
return resp
|
||||
}
|
||||
|
||||
// copyFloatPtr returns a fresh pointer to the same value, or nil. Keeps
|
||||
// stored models and API payloads from aliasing each other's rate fields.
|
||||
func copyFloatPtr(v *float64) *float64 {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
out := *v
|
||||
return &out
|
||||
}
|
||||
|
||||
// Copy returns a deep copy of the provider.
|
||||
func (p *Provider) Copy() *Provider {
|
||||
clone := *p
|
||||
if p.Models != nil {
|
||||
clone.Models = append([]ProviderModel(nil), p.Models...)
|
||||
clone.Models = make([]ProviderModel, len(p.Models))
|
||||
for i, m := range p.Models {
|
||||
m.CachedInputPer1k = copyFloatPtr(m.CachedInputPer1k)
|
||||
m.CacheReadPer1k = copyFloatPtr(m.CacheReadPer1k)
|
||||
m.CacheCreationPer1k = copyFloatPtr(m.CacheCreationPer1k)
|
||||
clone.Models[i] = m
|
||||
}
|
||||
}
|
||||
if p.ExtraValues != nil {
|
||||
clone.ExtraValues = make(map[string]string, len(p.ExtraValues))
|
||||
|
||||
@@ -103,6 +103,12 @@ func TestSynthesizedService_WireShape(t *testing.T) {
|
||||
|
||||
assert.Equal(t, middlewareIDCostMeter, mws[6].GetId(), "seventh middleware id")
|
||||
assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[6].GetSlot(), "cost meter slot")
|
||||
var costCfg costMeterConfig
|
||||
require.NoError(t, json.Unmarshal(mws[6].GetConfigJson(), &costCfg), "cost meter config JSON must decode from the wire")
|
||||
require.NotNil(t, costCfg.Pricing, "the pricing table must travel on the wire — the proxy has no embedded price list to fall back to")
|
||||
assert.NotEmpty(t, costCfg.Pricing.Defaults["openai"], "default table rides in every mapping")
|
||||
assert.NotEmpty(t, costCfg.Pricing.Defaults["anthropic"], "default table covers all surfaces")
|
||||
assert.NotEmpty(t, costCfg.Pricing.Defaults["bedrock"], "default table covers all surfaces")
|
||||
|
||||
assert.Equal(t, middlewareIDLLMResponseParser, mws[7].GetId(), "eighth middleware id")
|
||||
assert.Equal(t, proto.MiddlewareSlot_MIDDLEWARE_SLOT_ON_RESPONSE, mws[7].GetSlot(), "response parser slot")
|
||||
|
||||
@@ -55,6 +55,8 @@ type Config struct {
|
||||
|
||||
ReverseProxy ReverseProxy
|
||||
|
||||
AgentNetwork AgentNetwork
|
||||
|
||||
// disable default all-to-all policy
|
||||
DisableDefaultPolicy bool
|
||||
|
||||
@@ -185,6 +187,24 @@ type StoreConfig struct {
|
||||
Engine types.Engine
|
||||
}
|
||||
|
||||
// AgentNetwork contains agent-network (LLM gateway) configuration.
|
||||
type AgentNetwork struct {
|
||||
// PricingDefaultsFile is the path to the YAML file holding the default
|
||||
// LLM pricing table (defaults_llm_pricing.yaml). Empty falls back to
|
||||
// probing <Datadir>/defaults_llm_pricing.yaml; with no file present the
|
||||
// compiled-in defaults serve. Schema: surface ("openai"/"anthropic"/
|
||||
// "bedrock") -> model -> rates in USD per 1k tokens (input_per_1k,
|
||||
// output_per_1k, and the optional cached_input_per_1k /
|
||||
// cache_read_per_1k / cache_creation_per_1k). File entries replace the
|
||||
// compiled-in entry for the same surface+model whole; everything else
|
||||
// keeps the compiled-in rates. The file is re-read periodically (mtime
|
||||
// poll), and the live table feeds both the synthesizer (what proxies
|
||||
// bill with) and the dashboard's catalog endpoint (what model rows
|
||||
// prefill with). An explicitly configured path that fails to load
|
||||
// fails startup; runtime reload errors keep the previous table.
|
||||
PricingDefaultsFile string
|
||||
}
|
||||
|
||||
// ReverseProxy contains reverse proxy configuration in front of management.
|
||||
type ReverseProxy struct {
|
||||
// TrustedHTTPProxies represents a list of trusted HTTP proxies by their IP prefixes.
|
||||
|
||||
@@ -50,7 +50,7 @@ func ToComponentSyncResponse(
|
||||
// TODO (dmitri) consider using invariants?
|
||||
//
|
||||
enableSSH := computeSSHEnabledForPeer(components, peer)
|
||||
peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH)
|
||||
peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH, components.ForceRoutingPeerDNSResolution)
|
||||
|
||||
includeIPv6 := peer.SupportsIPv6() && peer.IPv6.IsValid()
|
||||
useSourcePrefixes := peer.SupportsSourcePrefixes()
|
||||
|
||||
@@ -119,7 +119,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken
|
||||
return nbConfig
|
||||
}
|
||||
|
||||
func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool) *proto.PeerConfig {
|
||||
func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig {
|
||||
netmask, _ := network.Net.Mask.Size()
|
||||
fqdn := peer.FQDN(dnsName)
|
||||
|
||||
@@ -135,7 +135,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, set
|
||||
Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask),
|
||||
SshConfig: sshConfig,
|
||||
Fqdn: fqdn,
|
||||
RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled,
|
||||
RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled || peer.ProxyMeta.Embedded || forceRoutingPeerDNS,
|
||||
LazyConnectionEnabled: settings.LazyConnectionEnabled,
|
||||
AutoUpdate: &proto.AutoUpdateSettings{
|
||||
Version: settings.AutoUpdateVersion,
|
||||
@@ -162,12 +162,12 @@ func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nb
|
||||
useSourcePrefixes := peer.SupportsSourcePrefixes()
|
||||
|
||||
response := &proto.SyncResponse{
|
||||
PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH),
|
||||
PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution),
|
||||
NetworkMap: &proto.NetworkMap{
|
||||
Serial: networkMap.Network.CurrentSerial(),
|
||||
Routes: networkmap.ToProtocolRoutes(networkMap.Routes),
|
||||
DNSConfig: networkmap.ToProtocolDNSConfig(networkMap.DNSConfig, dnsCache, dnsFwdPort),
|
||||
PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH),
|
||||
PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution),
|
||||
},
|
||||
Checks: toProtocolChecks(ctx, checks),
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package grpc
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"testing"
|
||||
@@ -14,6 +15,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/networkmap"
|
||||
)
|
||||
@@ -301,3 +303,35 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) {
|
||||
assert.True(t, nbCfg.Metrics.Enabled, "metrics flag should carry the settings value")
|
||||
})
|
||||
}
|
||||
|
||||
func TestToPeerConfig_RoutingPeerDNSResolution(t *testing.T) {
|
||||
network := &types.Network{Net: net.IPNet{IP: net.IPv4(100, 0, 0, 0), Mask: net.CIDRMask(8, 32)}}
|
||||
|
||||
newPeer := func(embedded bool) *nbpeer.Peer {
|
||||
p := &nbpeer.Peer{IP: netip.MustParseAddr("100.0.0.1")}
|
||||
p.ProxyMeta.Embedded = embedded
|
||||
return p
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
globalFlag bool
|
||||
embedded bool
|
||||
forceParam bool
|
||||
wantEnabled bool
|
||||
}{
|
||||
{name: "global off, regular peer, no force", wantEnabled: false},
|
||||
{name: "global on wins", globalFlag: true, wantEnabled: true},
|
||||
{name: "embedded proxy peer forced", embedded: true, wantEnabled: true},
|
||||
{name: "routing peer forced via param", forceParam: true, wantEnabled: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
settings := &types.Settings{RoutingPeerDNSResolutionEnabled: tt.globalFlag}
|
||||
cfg := toPeerConfig(newPeer(tt.embedded), network, "netbird.selfhosted", settings, nil, nil, false, tt.forceParam)
|
||||
assert.Equal(t, tt.wantEnabled, cfg.RoutingPeerDnsResolutionEnabled,
|
||||
"RoutingPeerDnsResolutionEnabled should reflect global || embedded || forced")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -285,6 +285,7 @@ func (s *ProxyServiceServer) CheckLLMPolicyLimits(ctx context.Context, req *prot
|
||||
UserID: req.GetUserId(),
|
||||
GroupIDs: req.GetGroupIds(),
|
||||
ProviderID: req.GetProviderId(),
|
||||
Model: req.GetModel(),
|
||||
})
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("select policy for request: %v", err)
|
||||
|
||||
138
management/internals/shared/grpc/proxy_llm_policy_limits_test.go
Normal file
138
management/internals/shared/grpc/proxy_llm_policy_limits_test.go
Normal file
@@ -0,0 +1,138 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
grpcstatus "google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// fakeAgentNetworkLimits records the PolicySelectionInput it was invoked with
|
||||
// and returns a pre-programmed result, so tests can assert what the handler
|
||||
// forwards to the selector.
|
||||
type fakeAgentNetworkLimits struct {
|
||||
gotInput agentnetwork.PolicySelectionInput
|
||||
result *agentnetwork.PolicySelectionResult
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeAgentNetworkLimits) SelectPolicyForRequest(_ context.Context, in agentnetwork.PolicySelectionInput) (*agentnetwork.PolicySelectionResult, error) {
|
||||
f.gotInput = in
|
||||
if f.err != nil {
|
||||
return nil, f.err
|
||||
}
|
||||
return f.result, nil
|
||||
}
|
||||
|
||||
func (f *fakeAgentNetworkLimits) RecordUsage(_ context.Context, _ agentnetwork.RecordUsageInput) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// TestCheckLLMPolicyLimits_ForwardsModelToSelector proves the wiring added here:
|
||||
// the model the proxy extracted must reach the selector's Model unchanged,
|
||||
// alongside the account/user/group/provider fields.
|
||||
func TestCheckLLMPolicyLimits_ForwardsModelToSelector(t *testing.T) {
|
||||
fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{Allow: true, SelectedPolicyID: "pol-1"}}
|
||||
s := &ProxyServiceServer{}
|
||||
s.SetAgentNetworkLimitsService(fake)
|
||||
|
||||
req := &proto.CheckLLMPolicyLimitsRequest{
|
||||
AccountId: "acc-1",
|
||||
UserId: "user-1",
|
||||
GroupIds: []string{"grp-a", "grp-b"},
|
||||
ProviderId: "prov-1",
|
||||
Model: "claude-opus-4",
|
||||
}
|
||||
|
||||
resp, err := s.CheckLLMPolicyLimits(context.Background(), req)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
|
||||
assert.Equal(t, "acc-1", fake.gotInput.AccountID)
|
||||
assert.Equal(t, "user-1", fake.gotInput.UserID)
|
||||
assert.Equal(t, []string{"grp-a", "grp-b"}, fake.gotInput.GroupIDs)
|
||||
assert.Equal(t, "prov-1", fake.gotInput.ProviderID)
|
||||
assert.Equal(t, "claude-opus-4", fake.gotInput.Model,
|
||||
"the request's model must be forwarded to the selector")
|
||||
}
|
||||
|
||||
// TestCheckLLMPolicyLimits_EmptyModelForwardedAsEmpty proves an undetermined
|
||||
// model (empty string) is forwarded as-is; the selector decides how to treat it.
|
||||
func TestCheckLLMPolicyLimits_EmptyModelForwardedAsEmpty(t *testing.T) {
|
||||
fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{Allow: true}}
|
||||
s := &ProxyServiceServer{}
|
||||
s.SetAgentNetworkLimitsService(fake)
|
||||
|
||||
req := &proto.CheckLLMPolicyLimitsRequest{
|
||||
AccountId: "acc-1",
|
||||
ProviderId: "prov-1",
|
||||
}
|
||||
|
||||
_, err := s.CheckLLMPolicyLimits(context.Background(), req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "", fake.gotInput.Model, "an absent model must be forwarded as empty, not fabricated")
|
||||
}
|
||||
|
||||
// TestCheckLLMPolicyLimits_DenyResponseCarriesModelBlockedCode proves the deny
|
||||
// envelope surfaces the model-allowlist deny code + reason through the response.
|
||||
func TestCheckLLMPolicyLimits_DenyResponseCarriesModelBlockedCode(t *testing.T) {
|
||||
fake := &fakeAgentNetworkLimits{result: &agentnetwork.PolicySelectionResult{
|
||||
Allow: false,
|
||||
DenyCode: "llm_policy.model_blocked",
|
||||
DenyReason: `model "claude-opus-4" is not permitted by any applicable policy allowlist`,
|
||||
}}
|
||||
s := &ProxyServiceServer{}
|
||||
s.SetAgentNetworkLimitsService(fake)
|
||||
|
||||
resp, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{
|
||||
AccountId: "acc-1",
|
||||
ProviderId: "prov-1",
|
||||
Model: "claude-opus-4",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
assert.Equal(t, "deny", resp.Decision)
|
||||
assert.Equal(t, "llm_policy.model_blocked", resp.DenyCode)
|
||||
assert.NotEmpty(t, resp.DenyReason)
|
||||
assert.Empty(t, resp.SelectedPolicyId, "a denied request must carry no selected policy")
|
||||
}
|
||||
|
||||
// TestCheckLLMPolicyLimits_SelectorErrorSurfacesAsInternal proves a selector
|
||||
// failure surfaces as an Internal gRPC error rather than a silent allow.
|
||||
func TestCheckLLMPolicyLimits_SelectorErrorSurfacesAsInternal(t *testing.T) {
|
||||
fake := &fakeAgentNetworkLimits{err: errors.New("boom")}
|
||||
s := &ProxyServiceServer{}
|
||||
s.SetAgentNetworkLimitsService(fake)
|
||||
|
||||
_, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{
|
||||
AccountId: "acc-1",
|
||||
ProviderId: "prov-1",
|
||||
Model: "gpt-4o",
|
||||
})
|
||||
require.Error(t, err)
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.Internal, st.Code(), "selector errors must never fail open on the hot path")
|
||||
}
|
||||
|
||||
// TestCheckLLMPolicyLimits_UnconfiguredServiceReturnsUnimplemented locks the
|
||||
// fallback: with no limits service wired the RPC returns Unimplemented.
|
||||
func TestCheckLLMPolicyLimits_UnconfiguredServiceReturnsUnimplemented(t *testing.T) {
|
||||
s := &ProxyServiceServer{}
|
||||
|
||||
_, err := s.CheckLLMPolicyLimits(context.Background(), &proto.CheckLLMPolicyLimitsRequest{
|
||||
AccountId: "acc-1",
|
||||
ProviderId: "prov-1",
|
||||
})
|
||||
require.Error(t, err)
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.Unimplemented, st.Code())
|
||||
}
|
||||
@@ -921,7 +921,7 @@ func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, ne
|
||||
// if peer has reached this point then it has logged in
|
||||
loginResp := &proto.LoginResponse{
|
||||
NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, settings),
|
||||
PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH),
|
||||
PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false),
|
||||
Checks: toProtocolChecks(ctx, postureChecks),
|
||||
}
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"net"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
@@ -2258,117 +2259,30 @@ func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*p
|
||||
return checks, nil
|
||||
}
|
||||
|
||||
func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) {
|
||||
const serviceQuery = `SELECT id, account_id, name, domain, enabled, auth,
|
||||
meta_created_at, meta_certificate_issued_at, meta_status, proxy_cluster,
|
||||
pass_host_header, rewrite_redirects, session_private_key, session_public_key,
|
||||
mode, listen_port, port_auto_assigned, source, source_peer, terminated,
|
||||
private, access_groups
|
||||
FROM services WHERE account_id = $1`
|
||||
// serviceSelectColumns and targetSelectColumns are the column lists the Postgres
|
||||
// pgx read path scans. They must stay in sync with the rpservice.Service and
|
||||
// rpservice.Target gorm models; TestPgxServiceColumnsMatchGorm enforces this.
|
||||
const serviceSelectColumns = `id, account_id, name, domain, enabled, auth, restrictions,
|
||||
meta_created_at, meta_certificate_issued_at, meta_last_renewed_at, meta_status, proxy_cluster,
|
||||
pass_host_header, rewrite_redirects, session_private_key, session_public_key,
|
||||
mode, listen_port, port_auto_assigned, source, source_peer, terminated,
|
||||
private, access_groups`
|
||||
|
||||
const targetsQuery = `SELECT id, account_id, service_id, path, host, port, protocol,
|
||||
target_id, target_type, enabled
|
||||
FROM targets WHERE service_id = ANY($1)`
|
||||
const targetSelectColumns = `id, account_id, service_id, path, host, port, protocol,
|
||||
target_id, target_type, enabled, proxy_protocol,
|
||||
skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers,
|
||||
direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes,
|
||||
capture_content_types, agent_network, disable_access_log`
|
||||
|
||||
func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) {
|
||||
const serviceQuery = `SELECT ` + serviceSelectColumns + ` FROM services WHERE account_id = $1`
|
||||
|
||||
serviceRows, err := s.pool.Query(ctx, serviceQuery, accountID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
services, err := pgx.CollectRows(serviceRows, func(row pgx.CollectableRow) (*rpservice.Service, error) {
|
||||
var s rpservice.Service
|
||||
var auth []byte
|
||||
var accessGroups []byte
|
||||
var createdAt, certIssuedAt sql.NullTime
|
||||
var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString
|
||||
var mode, source, sourcePeer sql.NullString
|
||||
var terminated, portAutoAssigned, private sql.NullBool
|
||||
var listenPort sql.NullInt64
|
||||
err := row.Scan(
|
||||
&s.ID,
|
||||
&s.AccountID,
|
||||
&s.Name,
|
||||
&s.Domain,
|
||||
&s.Enabled,
|
||||
&auth,
|
||||
&createdAt,
|
||||
&certIssuedAt,
|
||||
&status,
|
||||
&proxyCluster,
|
||||
&s.PassHostHeader,
|
||||
&s.RewriteRedirects,
|
||||
&sessionPrivateKey,
|
||||
&sessionPublicKey,
|
||||
&mode,
|
||||
&listenPort,
|
||||
&portAutoAssigned,
|
||||
&source,
|
||||
&sourcePeer,
|
||||
&terminated,
|
||||
&private,
|
||||
&accessGroups,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if auth != nil {
|
||||
if err := json.Unmarshal(auth, &s.Auth); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if len(accessGroups) > 0 {
|
||||
if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal access_groups: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if private.Valid {
|
||||
s.Private = private.Bool
|
||||
}
|
||||
|
||||
s.Meta = rpservice.Meta{}
|
||||
if createdAt.Valid {
|
||||
s.Meta.CreatedAt = createdAt.Time
|
||||
}
|
||||
if certIssuedAt.Valid {
|
||||
t := certIssuedAt.Time
|
||||
s.Meta.CertificateIssuedAt = &t
|
||||
}
|
||||
if status.Valid {
|
||||
s.Meta.Status = status.String
|
||||
}
|
||||
if proxyCluster.Valid {
|
||||
s.ProxyCluster = proxyCluster.String
|
||||
}
|
||||
if sessionPrivateKey.Valid {
|
||||
s.SessionPrivateKey = sessionPrivateKey.String
|
||||
}
|
||||
if sessionPublicKey.Valid {
|
||||
s.SessionPublicKey = sessionPublicKey.String
|
||||
}
|
||||
if mode.Valid {
|
||||
s.Mode = mode.String
|
||||
}
|
||||
if source.Valid {
|
||||
s.Source = source.String
|
||||
}
|
||||
if sourcePeer.Valid {
|
||||
s.SourcePeer = sourcePeer.String
|
||||
}
|
||||
if terminated.Valid {
|
||||
s.Terminated = terminated.Bool
|
||||
}
|
||||
if portAutoAssigned.Valid {
|
||||
s.PortAutoAssigned = portAutoAssigned.Bool
|
||||
}
|
||||
if listenPort.Valid {
|
||||
s.ListenPort = uint16(listenPort.Int64)
|
||||
}
|
||||
s.Targets = []*rpservice.Target{}
|
||||
return &s, nil
|
||||
})
|
||||
services, err := pgx.CollectRows(serviceRows, scanService)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -2379,39 +2293,12 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv
|
||||
|
||||
serviceIDs := make([]string, len(services))
|
||||
serviceMap := make(map[string]*rpservice.Service)
|
||||
for i, s := range services {
|
||||
serviceIDs[i] = s.ID
|
||||
serviceMap[s.ID] = s
|
||||
for i, svc := range services {
|
||||
serviceIDs[i] = svc.ID
|
||||
serviceMap[svc.ID] = svc
|
||||
}
|
||||
|
||||
targetRows, err := s.pool.Query(ctx, targetsQuery, serviceIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targets, err := pgx.CollectRows(targetRows, func(row pgx.CollectableRow) (*rpservice.Target, error) {
|
||||
var t rpservice.Target
|
||||
var path sql.NullString
|
||||
err := row.Scan(
|
||||
&t.ID,
|
||||
&t.AccountID,
|
||||
&t.ServiceID,
|
||||
&path,
|
||||
&t.Host,
|
||||
&t.Port,
|
||||
&t.Protocol,
|
||||
&t.TargetId,
|
||||
&t.TargetType,
|
||||
&t.Enabled,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if path.Valid {
|
||||
t.Path = &path.String
|
||||
}
|
||||
return &t, nil
|
||||
})
|
||||
targets, err := s.getServiceTargets(ctx, serviceIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -2425,6 +2312,201 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv
|
||||
return services, nil
|
||||
}
|
||||
|
||||
func scanService(row pgx.CollectableRow) (*rpservice.Service, error) {
|
||||
var s rpservice.Service
|
||||
var auth []byte
|
||||
var restrictions []byte
|
||||
var accessGroups []byte
|
||||
var createdAt, certIssuedAt, lastRenewedAt sql.NullTime
|
||||
var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString
|
||||
var mode, source, sourcePeer sql.NullString
|
||||
var terminated, portAutoAssigned, private sql.NullBool
|
||||
var listenPort sql.NullInt64
|
||||
err := row.Scan(
|
||||
&s.ID,
|
||||
&s.AccountID,
|
||||
&s.Name,
|
||||
&s.Domain,
|
||||
&s.Enabled,
|
||||
&auth,
|
||||
&restrictions,
|
||||
&createdAt,
|
||||
&certIssuedAt,
|
||||
&lastRenewedAt,
|
||||
&status,
|
||||
&proxyCluster,
|
||||
&s.PassHostHeader,
|
||||
&s.RewriteRedirects,
|
||||
&sessionPrivateKey,
|
||||
&sessionPublicKey,
|
||||
&mode,
|
||||
&listenPort,
|
||||
&portAutoAssigned,
|
||||
&source,
|
||||
&sourcePeer,
|
||||
&terminated,
|
||||
&private,
|
||||
&accessGroups,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if auth != nil {
|
||||
if err := json.Unmarshal(auth, &s.Auth); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if len(restrictions) > 0 {
|
||||
if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal restrictions: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if len(accessGroups) > 0 {
|
||||
if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal access_groups: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if private.Valid {
|
||||
s.Private = private.Bool
|
||||
}
|
||||
|
||||
s.Meta = serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt, status)
|
||||
if proxyCluster.Valid {
|
||||
s.ProxyCluster = proxyCluster.String
|
||||
}
|
||||
if sessionPrivateKey.Valid {
|
||||
s.SessionPrivateKey = sessionPrivateKey.String
|
||||
}
|
||||
if sessionPublicKey.Valid {
|
||||
s.SessionPublicKey = sessionPublicKey.String
|
||||
}
|
||||
if mode.Valid {
|
||||
s.Mode = mode.String
|
||||
}
|
||||
if source.Valid {
|
||||
s.Source = source.String
|
||||
}
|
||||
if sourcePeer.Valid {
|
||||
s.SourcePeer = sourcePeer.String
|
||||
}
|
||||
if terminated.Valid {
|
||||
s.Terminated = terminated.Bool
|
||||
}
|
||||
if portAutoAssigned.Valid {
|
||||
s.PortAutoAssigned = portAutoAssigned.Bool
|
||||
}
|
||||
if listenPort.Valid {
|
||||
if listenPort.Int64 < 0 || listenPort.Int64 > math.MaxUint16 {
|
||||
return nil, fmt.Errorf("listen_port %d out of range", listenPort.Int64)
|
||||
}
|
||||
s.ListenPort = uint16(listenPort.Int64)
|
||||
}
|
||||
s.Targets = []*rpservice.Target{}
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
func serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt sql.NullTime, status sql.NullString) rpservice.Meta {
|
||||
meta := rpservice.Meta{}
|
||||
if createdAt.Valid {
|
||||
meta.CreatedAt = createdAt.Time
|
||||
}
|
||||
if certIssuedAt.Valid {
|
||||
t := certIssuedAt.Time
|
||||
meta.CertificateIssuedAt = &t
|
||||
}
|
||||
if lastRenewedAt.Valid {
|
||||
t := lastRenewedAt.Time
|
||||
meta.LastRenewedAt = &t
|
||||
}
|
||||
if status.Valid {
|
||||
meta.Status = status.String
|
||||
}
|
||||
return meta
|
||||
}
|
||||
|
||||
func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) {
|
||||
const targetsQuery = `SELECT ` + targetSelectColumns + ` FROM targets WHERE service_id = ANY($1)`
|
||||
|
||||
rows, err := s.pool.Query(ctx, targetsQuery, serviceIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return pgx.CollectRows(rows, scanTarget)
|
||||
}
|
||||
|
||||
func scanTarget(row pgx.CollectableRow) (*rpservice.Target, error) {
|
||||
var t rpservice.Target
|
||||
var path sql.NullString
|
||||
var pathRewrite sql.NullString
|
||||
var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool
|
||||
var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64
|
||||
var customHeaders, middlewares, captureContentTypes []byte
|
||||
err := row.Scan(
|
||||
&t.ID,
|
||||
&t.AccountID,
|
||||
&t.ServiceID,
|
||||
&path,
|
||||
&t.Host,
|
||||
&t.Port,
|
||||
&t.Protocol,
|
||||
&t.TargetId,
|
||||
&t.TargetType,
|
||||
&t.Enabled,
|
||||
&proxyProtocol,
|
||||
&skipTLSVerify,
|
||||
&requestTimeout,
|
||||
&sessionIdleTimeout,
|
||||
&pathRewrite,
|
||||
&customHeaders,
|
||||
&directUpstream,
|
||||
&middlewares,
|
||||
&captureMaxRequestBytes,
|
||||
&captureMaxResponseBytes,
|
||||
&captureContentTypes,
|
||||
&agentNetwork,
|
||||
&disableAccessLog,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if path.Valid {
|
||||
t.Path = &path.String
|
||||
}
|
||||
|
||||
t.ProxyProtocol = proxyProtocol.Bool
|
||||
t.Options.SkipTLSVerify = skipTLSVerify.Bool
|
||||
t.Options.RequestTimeout = time.Duration(requestTimeout.Int64)
|
||||
t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64)
|
||||
t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String)
|
||||
t.Options.DirectUpstream = directUpstream.Bool
|
||||
t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64
|
||||
t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64
|
||||
t.Options.AgentNetwork = agentNetwork.Bool
|
||||
t.Options.DisableAccessLog = disableAccessLog.Bool
|
||||
|
||||
if len(customHeaders) > 0 {
|
||||
if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal custom_headers: %w", err)
|
||||
}
|
||||
}
|
||||
if len(middlewares) > 0 {
|
||||
if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal middlewares: %w", err)
|
||||
}
|
||||
}
|
||||
if len(captureContentTypes) > 0 {
|
||||
if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal capture_content_types: %w", err)
|
||||
}
|
||||
}
|
||||
return &t, nil
|
||||
}
|
||||
|
||||
func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) {
|
||||
const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1`
|
||||
rows, err := s.pool.Query(ctx, query, accountID)
|
||||
|
||||
74
management/server/store/sql_store_pgx_parity_test.go
Normal file
74
management/server/store/sql_store_pgx_parity_test.go
Normal file
@@ -0,0 +1,74 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm/schema"
|
||||
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
)
|
||||
|
||||
// TestPgxServiceColumnsMatchGorm guards the Postgres pgx read path against
|
||||
// drifting from the gorm model. The SQLite/MySQL gorm path loads rows by struct,
|
||||
// so a new column on a model is picked up automatically, but the hand-written
|
||||
// pgx SELECT in sql_store.go must be updated by hand. This test fails when a
|
||||
// gorm column is missing from the pgx column list, which otherwise silently
|
||||
// returns zero-valued on Postgres with no compile error.
|
||||
func TestPgxServiceColumnsMatchGorm(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
model any
|
||||
selectColumns string
|
||||
// excluded lists gorm columns intentionally not loaded by the pgx path.
|
||||
excluded map[string]struct{}
|
||||
}{
|
||||
{
|
||||
name: "service",
|
||||
model: &rpservice.Service{},
|
||||
selectColumns: serviceSelectColumns,
|
||||
},
|
||||
{
|
||||
name: "target",
|
||||
model: &rpservice.Target{},
|
||||
selectColumns: targetSelectColumns,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
selected := parseColumnList(tc.selectColumns)
|
||||
for _, col := range gormColumnNames(t, tc.model) {
|
||||
if _, ok := tc.excluded[col]; ok {
|
||||
continue
|
||||
}
|
||||
_, ok := selected[col]
|
||||
assert.Truef(t, ok,
|
||||
"gorm column %q is not read by the Postgres pgx SELECT; add it to %sSelectColumns in sql_store.go (or to the test's excluded set if it is intentionally not loaded)",
|
||||
col, tc.name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func parseColumnList(cols string) map[string]struct{} {
|
||||
set := make(map[string]struct{})
|
||||
for _, c := range strings.Split(cols, ",") {
|
||||
if c = strings.TrimSpace(c); c != "" {
|
||||
set[c] = struct{}{}
|
||||
}
|
||||
}
|
||||
return set
|
||||
}
|
||||
|
||||
// gormColumnNames returns the DB column names gorm would migrate for the model,
|
||||
// using the same default naming strategy the store configures.
|
||||
func gormColumnNames(t *testing.T, model any) []string {
|
||||
t.Helper()
|
||||
sch, err := schema.Parse(model, &sync.Map{}, schema.NamingStrategy{})
|
||||
require.NoError(t, err)
|
||||
return sch.DBNames
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"os"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -44,3 +45,91 @@ func TestSqlStore_GetAccount_PrivateServiceRoundtrip(t *testing.T) {
|
||||
assert.Equal(t, []string{"grp-admins", "grp-ops"}, got.AccessGroups)
|
||||
})
|
||||
}
|
||||
|
||||
// TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip guards the Postgres pgx
|
||||
// read path (getServices) against silently dropping columns present on the gorm
|
||||
// model. Before the fix these fields loaded correctly on SQLite but came back
|
||||
// zero-valued on Postgres because the hand-written SELECT and scan omitted them.
|
||||
func TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip(t *testing.T) {
|
||||
if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") {
|
||||
t.Skip("skip CI tests on darwin and windows")
|
||||
}
|
||||
|
||||
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||
ctx := context.Background()
|
||||
account := newAccountWithId(ctx, "account_svc_opts", "testuser", "")
|
||||
require.NoError(t, store.SaveAccount(ctx, account))
|
||||
|
||||
renewedAt := time.Now().UTC().Truncate(time.Second)
|
||||
targetPath := "/api"
|
||||
svc := &rpservice.Service{
|
||||
ID: "svc-opts",
|
||||
AccountID: account.Id,
|
||||
Name: "opts-svc",
|
||||
Domain: "opts.example",
|
||||
Enabled: true,
|
||||
Mode: rpservice.ModeHTTP,
|
||||
Restrictions: rpservice.AccessRestrictions{
|
||||
AllowedCIDRs: []string{"10.0.0.0/8"},
|
||||
BlockedCountries: []string{"XX"},
|
||||
CrowdSecMode: "block",
|
||||
},
|
||||
Meta: rpservice.Meta{
|
||||
LastRenewedAt: &renewedAt,
|
||||
},
|
||||
Targets: []*rpservice.Target{
|
||||
{
|
||||
AccountID: account.Id,
|
||||
ServiceID: "svc-opts",
|
||||
Path: &targetPath,
|
||||
Host: "backend.internal",
|
||||
Port: 8080,
|
||||
Protocol: "http",
|
||||
TargetId: "tgt-1",
|
||||
Enabled: true,
|
||||
ProxyProtocol: true,
|
||||
Options: rpservice.TargetOptions{
|
||||
SkipTLSVerify: true,
|
||||
RequestTimeout: 30 * time.Second,
|
||||
SessionIdleTimeout: 5 * time.Minute,
|
||||
PathRewrite: rpservice.PathRewritePreserve,
|
||||
CustomHeaders: map[string]string{"X-Foo": "bar"},
|
||||
DirectUpstream: true,
|
||||
CaptureMaxRequestBytes: 1024,
|
||||
CaptureMaxResponseBytes: 2048,
|
||||
CaptureContentTypes: []string{"application/json"},
|
||||
AgentNetwork: true,
|
||||
DisableAccessLog: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
require.NoError(t, store.CreateService(ctx, svc))
|
||||
|
||||
loaded, err := store.GetAccount(ctx, account.Id)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, loaded.Services, 1)
|
||||
|
||||
got := loaded.Services[0]
|
||||
assert.Equal(t, []string{"10.0.0.0/8"}, got.Restrictions.AllowedCIDRs, "restrictions allowed CIDRs")
|
||||
assert.Equal(t, []string{"XX"}, got.Restrictions.BlockedCountries, "restrictions blocked countries")
|
||||
assert.Equal(t, "block", got.Restrictions.CrowdSecMode, "restrictions crowdsec mode")
|
||||
require.NotNil(t, got.Meta.LastRenewedAt, "meta last renewed at")
|
||||
assert.WithinDuration(t, renewedAt, *got.Meta.LastRenewedAt, time.Second, "meta last renewed at")
|
||||
|
||||
require.Len(t, got.Targets, 1)
|
||||
tg := got.Targets[0]
|
||||
assert.True(t, tg.ProxyProtocol, "target proxy protocol")
|
||||
assert.True(t, tg.Options.SkipTLSVerify, "options skip TLS verify")
|
||||
assert.Equal(t, 30*time.Second, tg.Options.RequestTimeout, "options request timeout")
|
||||
assert.Equal(t, 5*time.Minute, tg.Options.SessionIdleTimeout, "options session idle timeout")
|
||||
assert.Equal(t, rpservice.PathRewritePreserve, tg.Options.PathRewrite, "options path rewrite")
|
||||
assert.Equal(t, map[string]string{"X-Foo": "bar"}, tg.Options.CustomHeaders, "options custom headers")
|
||||
assert.True(t, tg.Options.DirectUpstream, "options direct upstream")
|
||||
assert.Equal(t, int64(1024), tg.Options.CaptureMaxRequestBytes, "options capture max request bytes")
|
||||
assert.Equal(t, int64(2048), tg.Options.CaptureMaxResponseBytes, "options capture max response bytes")
|
||||
assert.Equal(t, []string{"application/json"}, tg.Options.CaptureContentTypes, "options capture content types")
|
||||
assert.True(t, tg.Options.AgentNetwork, "options agent network")
|
||||
assert.True(t, tg.Options.DisableAccessLog, "options disable access log")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1517,6 +1517,54 @@ func (a *Account) GetResourceRoutersMap() map[string]map[string]*routerTypes.Net
|
||||
return routers
|
||||
}
|
||||
|
||||
// forcesRoutingPeerDNSResolution reports whether the given peer must run
|
||||
// routing-peer DNS resolution regardless of the account-global
|
||||
// RoutingPeerDNSResolutionEnabled setting. It returns true when the peer is a
|
||||
// router for a domain network resource that is targeted by an enabled
|
||||
// reverse-proxy service, so the peer's DNS forwarder starts and can resolve
|
||||
// the target for the embedded proxy peers. Embedded proxy peers themselves are
|
||||
// handled at PeerConfig build time.
|
||||
func (a *Account) forcesRoutingPeerDNSResolution(peerID string, routers map[string]map[string]*routerTypes.NetworkRouter) bool {
|
||||
targeted := a.proxyTargetedDomainResourceIDs()
|
||||
if len(targeted) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, resource := range a.NetworkResources {
|
||||
if resource == nil || !resource.Enabled || resource.Type != resourceTypes.Domain {
|
||||
continue
|
||||
}
|
||||
if _, ok := targeted[resource.ID]; !ok {
|
||||
continue
|
||||
}
|
||||
if _, isRouter := routers[resource.NetworkID][peerID]; isRouter {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// proxyTargetedDomainResourceIDs returns the set of domain network resource IDs
|
||||
// targeted by an enabled, non-terminated reverse-proxy service.
|
||||
func (a *Account) proxyTargetedDomainResourceIDs() map[string]struct{} {
|
||||
ids := make(map[string]struct{})
|
||||
for _, svc := range a.Services {
|
||||
if svc == nil || !svc.Enabled || svc.Terminated {
|
||||
continue
|
||||
}
|
||||
for _, target := range svc.Targets {
|
||||
if target == nil || !target.Enabled {
|
||||
continue
|
||||
}
|
||||
if target.TargetType == service.TargetTypeDomain {
|
||||
ids[target.TargetId] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// getPoliciesSourcePeers collects all unique peers from the source groups defined in the given policies.
|
||||
func getPoliciesSourcePeers(policies []*Policy, groups map[string]*Group) map[string]struct{} {
|
||||
sourcePeers := make(map[string]struct{})
|
||||
|
||||
@@ -140,6 +140,8 @@ func (a *Account) GetPeerNetworkMapComponents(
|
||||
RouterPeers: make(map[string]*ComponentPeer),
|
||||
NetworkXIDToPublicID: make(map[string]string, len(a.Networks)),
|
||||
PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)),
|
||||
|
||||
ForceRoutingPeerDNSResolution: a.forcesRoutingPeerDNSResolution(peerID, routers),
|
||||
}
|
||||
for _, n := range a.Networks {
|
||||
if n != nil {
|
||||
|
||||
@@ -1751,3 +1751,71 @@ func hasPrivateAccessPolicy(account *Account, serviceID string) bool {
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func TestForcesRoutingPeerDNSResolution(t *testing.T) {
|
||||
buildAccountRes := func(serviceEnabled, targetEnabled, resourceEnabled bool, targetType service.TargetType, resType resourceTypes.NetworkResourceType) *Account {
|
||||
return &Account{
|
||||
Id: "accountID",
|
||||
Groups: map[string]*Group{
|
||||
"router-group": {ID: "router-group", Peers: []string{"router-peer-grp"}},
|
||||
},
|
||||
NetworkRouters: []*routerTypes.NetworkRouter{
|
||||
{ID: "r1", NetworkID: "net-1", AccountID: "accountID", Peer: "router-peer", Enabled: true},
|
||||
{ID: "r2", NetworkID: "net-1", AccountID: "accountID", PeerGroups: []string{"router-group"}, Enabled: true},
|
||||
},
|
||||
NetworkResources: []*resourceTypes.NetworkResource{
|
||||
{ID: "res-domain", AccountID: "accountID", NetworkID: "net-1", Type: resType, Domain: "example.org", Enabled: resourceEnabled},
|
||||
},
|
||||
Services: []*service.Service{
|
||||
{
|
||||
ID: "svc-1", AccountID: "accountID", Enabled: serviceEnabled,
|
||||
Targets: []*service.Target{
|
||||
{TargetId: "res-domain", TargetType: targetType, Enabled: targetEnabled},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
buildAccount := func(serviceEnabled, targetEnabled, resourceEnabled bool, targetType service.TargetType) *Account {
|
||||
return buildAccountRes(serviceEnabled, targetEnabled, resourceEnabled, targetType, resourceTypes.Domain)
|
||||
}
|
||||
|
||||
t.Run("router peer for RP-targeted domain resource is forced", func(t *testing.T) {
|
||||
account := buildAccount(true, true, true, service.TargetTypeDomain)
|
||||
routers := account.GetResourceRoutersMap()
|
||||
assert.True(t, account.forcesRoutingPeerDNSResolution("router-peer", routers), "direct router peer should be forced")
|
||||
assert.True(t, account.forcesRoutingPeerDNSResolution("router-peer-grp", routers), "group-member router peer should be forced")
|
||||
})
|
||||
|
||||
t.Run("non-router peer is not forced", func(t *testing.T) {
|
||||
account := buildAccount(true, true, true, service.TargetTypeDomain)
|
||||
assert.False(t, account.forcesRoutingPeerDNSResolution("other-peer", account.GetResourceRoutersMap()))
|
||||
})
|
||||
|
||||
t.Run("not forced when service disabled", func(t *testing.T) {
|
||||
account := buildAccount(false, true, true, service.TargetTypeDomain)
|
||||
assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()))
|
||||
})
|
||||
|
||||
t.Run("not forced when target disabled", func(t *testing.T) {
|
||||
account := buildAccount(true, false, true, service.TargetTypeDomain)
|
||||
assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()))
|
||||
})
|
||||
|
||||
t.Run("not forced when resource disabled", func(t *testing.T) {
|
||||
account := buildAccount(true, true, false, service.TargetTypeDomain)
|
||||
assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()))
|
||||
})
|
||||
|
||||
t.Run("not forced for non-domain target type", func(t *testing.T) {
|
||||
account := buildAccount(true, true, true, service.TargetTypePeer)
|
||||
assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()))
|
||||
})
|
||||
|
||||
t.Run("not forced when targeted resource is not a domain", func(t *testing.T) {
|
||||
account := buildAccountRes(true, true, true, service.TargetTypeDomain, resourceTypes.Host)
|
||||
assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()),
|
||||
"a domain target pointing at a non-domain resource must not force resolution")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -321,6 +321,10 @@ func (am *DefaultAccountManager) DeleteUser(ctx context.Context, accountID, init
|
||||
return err
|
||||
}
|
||||
|
||||
if targetUser.AccountID != accountID {
|
||||
return status.NewUserNotFoundError(targetUserID)
|
||||
}
|
||||
|
||||
if targetUser.Role == types.UserRoleOwner {
|
||||
return status.NewOwnerDeletePermissionError()
|
||||
}
|
||||
|
||||
@@ -802,6 +802,52 @@ func TestUser_DeleteUser_SelfDelete(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUser_DeleteUser_OtherAccount(t *testing.T) {
|
||||
testStore, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("Error when creating store: %s", err)
|
||||
}
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
account := newAccountWithId(context.Background(), mockAccountID, mockUserID, "", "", "", false)
|
||||
if err = testStore.SaveAccount(context.Background(), account); err != nil {
|
||||
t.Fatalf("Error when saving account: %s", err)
|
||||
}
|
||||
|
||||
otherAccount := newAccountWithId(context.Background(), "otherAccount", "otherOwner", "", "", "", false)
|
||||
otherAccount.Users["otherRegularUser"] = &types.User{
|
||||
Id: "otherRegularUser",
|
||||
AccountID: "otherAccount",
|
||||
Role: types.UserRoleUser,
|
||||
}
|
||||
otherAccount.Users["otherServiceUser"] = &types.User{
|
||||
Id: "otherServiceUser",
|
||||
AccountID: "otherAccount",
|
||||
Role: types.UserRoleUser,
|
||||
IsServiceUser: true,
|
||||
ServiceUserName: "otherServiceUser",
|
||||
}
|
||||
if err = testStore.SaveAccount(context.Background(), otherAccount); err != nil {
|
||||
t.Fatalf("Error when saving other account: %s", err)
|
||||
}
|
||||
|
||||
am := DefaultAccountManager{
|
||||
Store: testStore,
|
||||
eventStore: &activity.InMemoryEventStore{},
|
||||
permissionsManager: permissions.NewManager(testStore),
|
||||
}
|
||||
|
||||
for _, targetUserID := range []string{"otherRegularUser", "otherServiceUser"} {
|
||||
t.Run(targetUserID, func(t *testing.T) {
|
||||
err := am.DeleteUser(context.Background(), mockAccountID, mockUserID, targetUserID)
|
||||
assert.Equal(t, status.NewUserNotFoundError(targetUserID), err)
|
||||
|
||||
_, err = testStore.GetUserByUserID(context.Background(), store.LockingStrengthNone, targetUserID)
|
||||
assert.NoError(t, err, "user of another account must not be deleted")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUser_DeleteUser_regularUser(t *testing.T) {
|
||||
store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir())
|
||||
if err != nil {
|
||||
|
||||
@@ -3,20 +3,30 @@ package auth
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/sync/singleflight"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// tunnelCacheTTL caps how long a positive ValidateTunnelPeer result is
|
||||
// reused before re-fetching from management. 5 minutes balances freshness
|
||||
// against management load on busy mesh networks.
|
||||
// tunnelCacheTTL is the default cap on how long a positive ValidateTunnelPeer
|
||||
// result is reused before re-fetching from management. 5 minutes balances
|
||||
// freshness against management load on busy mesh networks. Override it with
|
||||
// envTunnelCacheTTL when an account needs authorization changes to take effect
|
||||
// sooner (at the cost of more ValidateTunnelPeer RPCs).
|
||||
const tunnelCacheTTL = 300 * time.Second
|
||||
|
||||
// envTunnelCacheTTL overrides tunnelCacheTTL. The value is a Go duration string
|
||||
// (e.g. "30s", "2m"); an unset, unparseable, or non-positive value keeps the
|
||||
// default.
|
||||
const envTunnelCacheTTL = "NB_PROXY_TUNNEL_CACHE_TTL"
|
||||
|
||||
// tunnelCachePerAccount caps the number of cached identities per account.
|
||||
// Bounded eviction avoids memory growth in pathological cases (huge peer
|
||||
// roster, brief request bursts) while staying generous for normal use.
|
||||
@@ -60,16 +70,35 @@ type accountBucket struct {
|
||||
order []tunnelCacheKey
|
||||
}
|
||||
|
||||
// newTunnelValidationCache constructs a cache with default TTL and bounds.
|
||||
// newTunnelValidationCache constructs a cache with the configured TTL
|
||||
// (envTunnelCacheTTL override or default) and default bounds.
|
||||
func newTunnelValidationCache() *tunnelValidationCache {
|
||||
return &tunnelValidationCache{
|
||||
entries: make(map[types.AccountID]*accountBucket),
|
||||
ttl: tunnelCacheTTL,
|
||||
ttl: tunnelCacheTTLFromEnv(),
|
||||
maxSize: tunnelCachePerAccount,
|
||||
now: time.Now,
|
||||
}
|
||||
}
|
||||
|
||||
// tunnelCacheTTLFromEnv returns the tunnel-cache TTL, honoring the
|
||||
// envTunnelCacheTTL override. The override must be a positive Go duration
|
||||
// string (e.g. "30s", "2m"); anything unset, unparseable, or non-positive
|
||||
// falls back to tunnelCacheTTL.
|
||||
func tunnelCacheTTLFromEnv() time.Duration {
|
||||
raw := strings.TrimSpace(os.Getenv(envTunnelCacheTTL))
|
||||
if raw == "" {
|
||||
return tunnelCacheTTL
|
||||
}
|
||||
d, err := time.ParseDuration(raw)
|
||||
if err != nil || d <= 0 {
|
||||
log.Warnf("ignoring invalid %s=%q (want a positive Go duration like 30s or 2m); using default %s",
|
||||
envTunnelCacheTTL, raw, tunnelCacheTTL)
|
||||
return tunnelCacheTTL
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// get returns a cached response for the key, or nil when missing or
|
||||
// expired. Expired entries are evicted lazily on read.
|
||||
func (c *tunnelValidationCache) get(key tunnelCacheKey) *proto.ValidateTunnelPeerResponse {
|
||||
|
||||
@@ -169,3 +169,32 @@ func TestTunnelCache_BoundedSizeEvictsOldest(t *testing.T) {
|
||||
assert.NotNil(t, cache.get(keys[1]), "second-newest must remain cached")
|
||||
assert.NotNil(t, cache.get(keys[2]), "newest must remain cached")
|
||||
}
|
||||
|
||||
func TestTunnelCacheTTLFromEnv(t *testing.T) {
|
||||
t.Run("unset uses default", func(t *testing.T) {
|
||||
t.Setenv(envTunnelCacheTTL, "")
|
||||
assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv())
|
||||
})
|
||||
t.Run("valid duration overrides", func(t *testing.T) {
|
||||
t.Setenv(envTunnelCacheTTL, "45s")
|
||||
assert.Equal(t, 45*time.Second, tunnelCacheTTLFromEnv())
|
||||
})
|
||||
t.Run("whitespace trimmed", func(t *testing.T) {
|
||||
t.Setenv(envTunnelCacheTTL, " 2m ")
|
||||
assert.Equal(t, 2*time.Minute, tunnelCacheTTLFromEnv())
|
||||
})
|
||||
t.Run("unparseable uses default", func(t *testing.T) {
|
||||
t.Setenv(envTunnelCacheTTL, "nonsense")
|
||||
assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv())
|
||||
})
|
||||
t.Run("non-positive uses default", func(t *testing.T) {
|
||||
t.Setenv(envTunnelCacheTTL, "0s")
|
||||
assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv())
|
||||
t.Setenv(envTunnelCacheTTL, "-30s")
|
||||
assert.Equal(t, tunnelCacheTTL, tunnelCacheTTLFromEnv())
|
||||
})
|
||||
t.Run("constructor honors override", func(t *testing.T) {
|
||||
t.Setenv(envTunnelCacheTTL, "90s")
|
||||
assert.Equal(t, 90*time.Second, newTunnelValidationCache().ttl)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,38 +0,0 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// bedrockRegionPrefixes are the cross-region inference-profile prefixes that
|
||||
// front a Bedrock model id (e.g. "eu.anthropic.claude-...").
|
||||
var bedrockRegionPrefixes = []string{"us.", "eu.", "apac.", "global."}
|
||||
|
||||
// bedrockVersionSuffix matches the trailing "-vN[:N]" or "-YYYYMMDD-vN[:N]"
|
||||
// version/throughput suffix of a Bedrock model id.
|
||||
var bedrockVersionSuffix = regexp.MustCompile(`-(\d{8}-)?v\d+(:\d+)?$`)
|
||||
|
||||
// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile
|
||||
// prefix, and the version/throughput suffix from a Bedrock model id so it
|
||||
// matches the catalog/pricing key, e.g.
|
||||
// "eu.anthropic.claude-sonnet-4-5-20250929-v1:0" -> "anthropic.claude-sonnet-4-5"
|
||||
// and the inference-profile ARN's last segment likewise. It is the single
|
||||
// source of truth shared by the request parser (which normalizes the request
|
||||
// model from the URL path) and the router (which normalizes the operator's
|
||||
// registered Bedrock model ids so both sides compare equal).
|
||||
func NormalizeBedrockModel(modelID string) string {
|
||||
m := modelID
|
||||
if strings.HasPrefix(m, "arn:") {
|
||||
if i := strings.LastIndex(m, "/"); i >= 0 {
|
||||
m = m[i+1:]
|
||||
}
|
||||
}
|
||||
for _, p := range bedrockRegionPrefixes {
|
||||
if strings.HasPrefix(m, p) {
|
||||
m = m[len(p):]
|
||||
break
|
||||
}
|
||||
}
|
||||
return bedrockVersionSuffix.ReplaceAllString(m, "")
|
||||
}
|
||||
@@ -1,59 +0,0 @@
|
||||
# Realistic-pricing starter for llm_observability. Drop this into the
|
||||
# directory you point the proxy at via --plugin-data-dir, then reference it
|
||||
# from the target's plugin config:
|
||||
#
|
||||
# plugins:
|
||||
# - id: llm_observability
|
||||
# enabled: true
|
||||
# params:
|
||||
# pricing_path: pricing.yaml
|
||||
#
|
||||
# Values are USD per 1_000 tokens. Public list prices drift; treat this as a
|
||||
# starting point and keep your production copy current.
|
||||
|
||||
openai:
|
||||
# GPT-5 family
|
||||
gpt-5:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
gpt-5-mini:
|
||||
input_per_1k: 0.00025
|
||||
output_per_1k: 0.002
|
||||
gpt-5-nano:
|
||||
input_per_1k: 0.00005
|
||||
output_per_1k: 0.0004
|
||||
gpt-5.4:
|
||||
input_per_1k: 0.00125
|
||||
output_per_1k: 0.01
|
||||
# GPT-4o family
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
gpt-4o-mini:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
# Embeddings
|
||||
text-embedding-3-large:
|
||||
input_per_1k: 0.00013
|
||||
output_per_1k: 0
|
||||
text-embedding-3-small:
|
||||
input_per_1k: 0.00002
|
||||
output_per_1k: 0
|
||||
|
||||
anthropic:
|
||||
# Claude 4.x family
|
||||
claude-opus-4-7:
|
||||
input_per_1k: 0.015
|
||||
output_per_1k: 0.075
|
||||
claude-sonnet-4-7:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
claude-sonnet-4-6:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
claude-haiku-4-5:
|
||||
input_per_1k: 0.0008
|
||||
output_per_1k: 0.004
|
||||
21
proxy/internal/llm/model.go
Normal file
21
proxy/internal/llm/model.go
Normal file
@@ -0,0 +1,21 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
sharedllm "github.com/netbirdio/netbird/shared/llm"
|
||||
)
|
||||
|
||||
// NormalizeBedrockModel strips an ARN wrapper, a cross-region inference-profile
|
||||
// prefix, and the version/throughput suffix from a Bedrock model id so it
|
||||
// matches the catalog/pricing key. Thin delegate to the shared implementation
|
||||
// (shared/llm), which management also uses at synthesis time so both sides of
|
||||
// the pricing / routing contract normalize identically.
|
||||
func NormalizeBedrockModel(modelID string) string {
|
||||
return sharedllm.NormalizeBedrockModel(modelID)
|
||||
}
|
||||
|
||||
// NormalizeVertexModel strips the "@version" suffix from a Vertex AI model id
|
||||
// so it matches the catalog/pricing key. Thin delegate to shared/llm, kept
|
||||
// beside NormalizeBedrockModel for the same contract reason.
|
||||
func NormalizeVertexModel(modelID string) string {
|
||||
return sharedllm.NormalizeVertexModel(modelID)
|
||||
}
|
||||
@@ -1,65 +0,0 @@
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestDefaultTable_FirstPartyModelCoverage guards the embedded defaults against
|
||||
// silent drift/gaps: every metered first-party model the management catalog
|
||||
// enumerates must resolve to a price, and a few rates that previously drifted
|
||||
// are pinned to their LiteLLM-validated values. Keep this list in step with the
|
||||
// catalog (management/server/agentnetwork/catalog) when adding models.
|
||||
func TestDefaultTable_FirstPartyModelCoverage(t *testing.T) {
|
||||
tbl := DefaultTable()
|
||||
require.NotNil(t, tbl, "embedded default pricing table must load")
|
||||
|
||||
mustPrice := map[string][]string{
|
||||
// openai parser covers openai_api, azure_openai_api, and mistral_api.
|
||||
"openai": {
|
||||
"gpt-5.5", "gpt-5.5-pro", "gpt-5.4", "gpt-5.4-mini", "gpt-5.4-nano",
|
||||
"gpt-5.3-codex", "gpt-5.3-chat-latest", "o4-mini",
|
||||
"gpt-4.1", "gpt-4.1-mini", "gpt-4.1-nano", "gpt-4o", "gpt-4o-mini",
|
||||
"gpt-4-turbo", "gpt-3.5-turbo", "gpt-35-turbo",
|
||||
"text-embedding-3-large", "text-embedding-3-small",
|
||||
"mistral-large-latest", "mistral-medium-3-5", "codestral-2508",
|
||||
"ministral-8b-latest", "mistral-embed",
|
||||
},
|
||||
"anthropic": {
|
||||
"claude-fable-5", "claude-opus-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6",
|
||||
"claude-opus-4-1", "claude-sonnet-4-6", "claude-sonnet-4-5", "claude-haiku-4-5",
|
||||
},
|
||||
// bedrock keys are the normalized ids the request parser emits.
|
||||
"bedrock": {
|
||||
"anthropic.claude-opus-5", "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6",
|
||||
"anthropic.claude-opus-4-1", "anthropic.claude-sonnet-4-6", "anthropic.claude-sonnet-4-5",
|
||||
"anthropic.claude-haiku-4-5", "meta.llama3-3-70b-instruct",
|
||||
"amazon.nova-pro", "amazon.nova-lite", "amazon.nova-micro", "amazon.nova-2-lite",
|
||||
},
|
||||
}
|
||||
for provider, models := range mustPrice {
|
||||
for _, m := range models {
|
||||
_, ok := tbl.Cost(provider, m, 1000, 1000, 0, 0)
|
||||
assert.True(t, ok, "%s/%s must be priced in the embedded defaults", provider, m)
|
||||
}
|
||||
}
|
||||
|
||||
// Pin per-direction rates independently (input-only then output-only) so a
|
||||
// swap or skew of input<->output that preserves the combined total is still
|
||||
// caught — these are rates that previously drifted or are easy to mis-enter.
|
||||
in, ok := tbl.Cost("openai", "gpt-5.4", 1000, 0, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.0025, in, 1e-9, "gpt-5.4 input = 0.0025 per 1k")
|
||||
out, ok := tbl.Cost("openai", "gpt-5.4", 0, 1000, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.015, out, 1e-9, "gpt-5.4 output = 0.015 per 1k")
|
||||
|
||||
in, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 1000, 0, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.003, in, 1e-9, "bedrock sonnet-4-5 input = 0.003 per 1k")
|
||||
out, ok = tbl.Cost("bedrock", "anthropic.claude-sonnet-4-5", 0, 1000, 0, 0)
|
||||
require.True(t, ok)
|
||||
assert.InDelta(t, 0.015, out, 1e-9, "bedrock sonnet-4-5 output = 0.015 per 1k")
|
||||
}
|
||||
@@ -1,102 +1,30 @@
|
||||
// Package pricing implements the embedded-default + override pricing table
|
||||
// shared by middleware that converts LLM token usage into a USD cost
|
||||
// estimate. The table is hot-reloadable from a basename under the proxy
|
||||
// data directory; missing override files keep the embedded defaults so
|
||||
// cost annotation works without operator action.
|
||||
// Package pricing implements the pricing table and cost formula the
|
||||
// cost_meter middleware uses to convert LLM token usage into a USD cost
|
||||
// estimate. The table's content arrives from the management server inside
|
||||
// cost_meter's middleware config (synthesized from the catalog plus the
|
||||
// operator's stored per-provider prices) — the proxy carries no embedded
|
||||
// price list. Price updates ride the ordinary mapping push: a chain
|
||||
// rebuild constructs a fresh table, so there is nothing to reload.
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
_ "embed"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"math"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
//go:embed defaults_pricing.yaml
|
||||
var defaultPricingYAML []byte
|
||||
|
||||
var (
|
||||
defaultTableOnce sync.Once
|
||||
defaultTablePtr *Table
|
||||
)
|
||||
|
||||
// DefaultTable returns the pricing table embedded in the binary. The result
|
||||
// is parsed once and shared; callers must not mutate the returned value.
|
||||
// Cost annotation works without any operator action because every loader
|
||||
// starts with this table.
|
||||
func DefaultTable() *Table {
|
||||
defaultTableOnce.Do(func() {
|
||||
t, err := parsePricingBytes(defaultPricingYAML)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("llmobs: embedded default pricing failed to parse: %v", err))
|
||||
}
|
||||
defaultTablePtr = t
|
||||
})
|
||||
return defaultTablePtr
|
||||
}
|
||||
|
||||
// mergeOver returns a new Table containing every entry from base, with any
|
||||
// matching entry from overlay replacing the base value. Either argument may
|
||||
// be nil. Result is a fresh allocation so callers can mutate / Store safely.
|
||||
func mergeOver(base, overlay *Table) *Table {
|
||||
if overlay == nil || len(overlay.entries) == 0 {
|
||||
return base
|
||||
}
|
||||
if base == nil || len(base.entries) == 0 {
|
||||
return overlay
|
||||
}
|
||||
out := make(map[string]map[string]Entry, len(base.entries))
|
||||
for provider, models := range base.entries {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, e := range models {
|
||||
inner[model] = e
|
||||
}
|
||||
out[provider] = inner
|
||||
}
|
||||
for provider, models := range overlay.entries {
|
||||
inner, ok := out[provider]
|
||||
if !ok {
|
||||
inner = make(map[string]Entry, len(models))
|
||||
out[provider] = inner
|
||||
}
|
||||
for model, e := range models {
|
||||
inner[model] = e
|
||||
}
|
||||
}
|
||||
return &Table{entries: out}
|
||||
}
|
||||
|
||||
// Entry is a single model's input and output pricing, expressed in USD per
|
||||
// 1000 tokens.
|
||||
//
|
||||
// CachedInputPer1K applies to OpenAI's cached prompt tokens, which are a
|
||||
// subset of input_tokens — when set, the cached portion is billed at this
|
||||
// rate and the non-cached remainder at InputPer1K. Zero means "no discount
|
||||
// configured", and cached tokens are billed at InputPer1K (matches current
|
||||
// behaviour where cached counts weren't extracted at all).
|
||||
// configured", and cached tokens are billed at InputPer1K.
|
||||
//
|
||||
// CacheReadPer1K and CacheCreationPer1K apply to Anthropic's two prompt-
|
||||
// cache fields, which are additive to input_tokens: cache_read is the
|
||||
// cheaper read-from-cache rate, cache_creation is the more expensive
|
||||
// write-to-cache rate. Zero means "no rate configured" and the
|
||||
// corresponding token bucket is billed at InputPer1K. This is more
|
||||
// accurate than today's behaviour, where Anthropic's cache tokens are
|
||||
// ignored and not charged at all.
|
||||
// corresponding token bucket is billed at InputPer1K.
|
||||
type Entry struct {
|
||||
InputPer1K float64
|
||||
OutputPer1K float64
|
||||
@@ -105,33 +33,102 @@ type Entry struct {
|
||||
CacheCreationPer1K float64
|
||||
}
|
||||
|
||||
// Table is a provider-to-model pricing lookup. Instances are immutable once
|
||||
// built and are swapped atomically by Loader.
|
||||
// EntryJSON is the wire shape of a pricing entry inside cost_meter's
|
||||
// middleware config. Field names are the management→proxy contract; the
|
||||
// management synthesizer marshals the same names (its pricing.Entry).
|
||||
type EntryJSON struct {
|
||||
InputPer1K float64 `json:"input_per_1k"`
|
||||
OutputPer1K float64 `json:"output_per_1k"`
|
||||
CachedInputPer1K float64 `json:"cached_input_per_1k"`
|
||||
CacheReadPer1K float64 `json:"cache_read_per_1k"`
|
||||
CacheCreationPer1K float64 `json:"cache_creation_per_1k"`
|
||||
}
|
||||
|
||||
// Table is a provider-surface-to-model pricing lookup. Instances are
|
||||
// immutable once built; a mapping update builds a whole new middleware
|
||||
// instance (and with it a new table) rather than mutating this one.
|
||||
type Table struct {
|
||||
entries map[string]map[string]Entry
|
||||
}
|
||||
|
||||
// NewEntries validates and converts a wire-shape map (surface-or-record ->
|
||||
// model -> rates) into the internal representation. Every rate must be a
|
||||
// finite, non-negative USD amount; a violation is returned as an error so
|
||||
// a corrupt config fails the chain build loudly instead of mispricing.
|
||||
// Management validates the same constraints at its API boundary, so this
|
||||
// is defense-in-depth. Nil input yields an empty (never-matching) map.
|
||||
func NewEntries(raw map[string]map[string]EntryJSON) (map[string]map[string]Entry, error) {
|
||||
out := make(map[string]map[string]Entry, len(raw))
|
||||
for outer, models := range raw {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, e := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input_per_1k": e.InputPer1K,
|
||||
"output_per_1k": e.OutputPer1K,
|
||||
"cached_input_per_1k": e.CachedInputPer1K,
|
||||
"cache_read_per_1k": e.CacheReadPer1K,
|
||||
"cache_creation_per_1k": e.CacheCreationPer1K,
|
||||
} {
|
||||
if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", outer, model, field, v)
|
||||
}
|
||||
}
|
||||
// EntryJSON and Entry are field-identical (tags aside), so a
|
||||
// direct conversion carries all five rates.
|
||||
inner[model] = Entry(e)
|
||||
}
|
||||
out[outer] = inner
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// NewTable builds an immutable Table from the wire-shape defaults map.
|
||||
// See NewEntries for validation semantics.
|
||||
func NewTable(raw map[string]map[string]EntryJSON) (*Table, error) {
|
||||
entries, err := NewEntries(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Table{entries: entries}, nil
|
||||
}
|
||||
|
||||
// Lookup returns the entry for the given provider surface and model.
|
||||
func (t *Table) Lookup(provider, model string) (Entry, bool) {
|
||||
if t == nil {
|
||||
return Entry{}, false
|
||||
}
|
||||
byModel, ok := t.entries[provider]
|
||||
if !ok {
|
||||
return Entry{}, false
|
||||
}
|
||||
e, ok := byModel[model]
|
||||
return e, ok
|
||||
}
|
||||
|
||||
// Has reports whether the provider/model pair is present in the table.
|
||||
func (t *Table) Has(provider, model string) bool {
|
||||
_, ok := t.Lookup(provider, model)
|
||||
return ok
|
||||
}
|
||||
|
||||
// Cost returns the estimated USD cost for the given token counts. ok is
|
||||
// false when the provider or model is not present in the table; the caller
|
||||
// can still emit token metrics with a model=unknown label.
|
||||
//
|
||||
// Provider-shape semantics for cached / cache-creation counts:
|
||||
//
|
||||
// - OpenAI: cachedInput is a SUBSET of inTokens. The cached portion is
|
||||
// billed at CachedInputPer1K (or InputPer1K when no override), and the
|
||||
// non-cached remainder of inTokens at InputPer1K. cacheCreation is
|
||||
// ignored (OpenAI has no analogue).
|
||||
// - Anthropic: cachedInput (cache_read) and cacheCreation are ADDITIVE to
|
||||
// inTokens. The three buckets are billed at CacheReadPer1K,
|
||||
// CacheCreationPer1K, and InputPer1K respectively, each falling back
|
||||
// to InputPer1K when the corresponding rate is zero.
|
||||
// - Other providers: cached and cacheCreation are ignored; cost is
|
||||
// inTokens*InputPer1K + outTokens*OutputPer1K.
|
||||
func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (float64, bool) {
|
||||
c, ok := t.Costs(provider, model, inTokens, outTokens, cachedInput, cacheCreation)
|
||||
return c.TotalUSD, ok
|
||||
}
|
||||
|
||||
// Costs returns the estimated USD cost split for the given token counts.
|
||||
// The provider surface selects the cache formula; see EntryCosts.
|
||||
func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) {
|
||||
entry, ok := t.Lookup(provider, model)
|
||||
if !ok {
|
||||
return Costs{}, false
|
||||
}
|
||||
return EntryCosts(entry, provider, inTokens, outTokens, cachedInput, cacheCreation), true
|
||||
}
|
||||
|
||||
// Costs is a per-request cost split. The four per-bucket fields are the base
|
||||
// of the breakdown — one per token bucket the provider bills separately — and
|
||||
// the two aggregates are derived from them:
|
||||
@@ -165,9 +162,25 @@ func newCosts(input, cachedInput, cacheCreation, output float64) Costs {
|
||||
}
|
||||
}
|
||||
|
||||
// Costs returns the estimated USD cost split for the given token counts, with
|
||||
// the same semantics as Cost.
|
||||
func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) {
|
||||
// EntryCosts computes the USD cost split for the given entry and token
|
||||
// counts. The surface (the llm.provider value the request parser stamps)
|
||||
// selects the cache formula; the entry may come from the surface-keyed
|
||||
// defaults table or from a per-provider-record override — the math is
|
||||
// identical either way.
|
||||
//
|
||||
// Provider-shape semantics for cached / cache-creation counts:
|
||||
//
|
||||
// - "openai": cachedInput is a SUBSET of inTokens. The cached portion is
|
||||
// billed at CachedInputPer1K (or InputPer1K when no override), and the
|
||||
// non-cached remainder of inTokens at InputPer1K. cacheCreation is
|
||||
// ignored (OpenAI has no analogue).
|
||||
// - "anthropic", "bedrock": cachedInput (cache_read) and cacheCreation are
|
||||
// ADDITIVE to inTokens. The three buckets are billed at CacheReadPer1K,
|
||||
// CacheCreationPer1K, and InputPer1K respectively, each falling back
|
||||
// to InputPer1K when the corresponding rate is zero.
|
||||
// - Other surfaces: cached and cacheCreation are ignored; cost is
|
||||
// inTokens*InputPer1K + outTokens*OutputPer1K.
|
||||
func EntryCosts(entry Entry, surface string, inTokens, outTokens, cachedInput, cacheCreation int64) Costs {
|
||||
// Clamp negatives to zero before any pricing math so a malformed
|
||||
// upstream count can never produce a negative cost.
|
||||
if inTokens < 0 {
|
||||
@@ -182,19 +195,8 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput,
|
||||
if cacheCreation < 0 {
|
||||
cacheCreation = 0
|
||||
}
|
||||
if t == nil {
|
||||
return Costs{}, false
|
||||
}
|
||||
byModel, ok := t.entries[provider]
|
||||
if !ok {
|
||||
return Costs{}, false
|
||||
}
|
||||
entry, ok := byModel[model]
|
||||
if !ok {
|
||||
return Costs{}, false
|
||||
}
|
||||
output := (float64(outTokens) / 1000.0) * entry.OutputPer1K
|
||||
switch provider {
|
||||
switch surface {
|
||||
case "openai":
|
||||
// cachedInput is a subset of inTokens; clamp so a malformed
|
||||
// upstream (cached > total) can't produce a negative remainder.
|
||||
@@ -208,7 +210,7 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput,
|
||||
}
|
||||
nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K
|
||||
cached := float64(clamped) / 1000.0 * cachedRate
|
||||
return newCosts(nonCached, cached, 0, output), true
|
||||
return newCosts(nonCached, cached, 0, output)
|
||||
case "anthropic", "bedrock":
|
||||
// Bedrock-Anthropic returns the same additive cache buckets as
|
||||
// first-party Anthropic; non-Anthropic Bedrock models simply report
|
||||
@@ -224,266 +226,9 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput,
|
||||
input := float64(inTokens) / 1000.0 * entry.InputPer1K
|
||||
read := float64(cachedInput) / 1000.0 * readRate
|
||||
create := float64(cacheCreation) / 1000.0 * createRate
|
||||
return newCosts(input, read, create, output), true
|
||||
return newCosts(input, read, create, output)
|
||||
default:
|
||||
input := float64(inTokens) / 1000.0 * entry.InputPer1K
|
||||
return newCosts(input, 0, 0, output), true
|
||||
return newCosts(input, 0, 0, output)
|
||||
}
|
||||
}
|
||||
|
||||
// Has reports whether the provider/model pair is present in the table.
|
||||
func (t *Table) Has(provider, model string) bool {
|
||||
if t == nil {
|
||||
return false
|
||||
}
|
||||
byModel, ok := t.entries[provider]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
_, ok = byModel[model]
|
||||
return ok
|
||||
}
|
||||
|
||||
// pricingFile mirrors the on-disk YAML schema. Keys are provider names; the
|
||||
// nested map keys are model names.
|
||||
type pricingFile map[string]map[string]struct {
|
||||
InputPer1K float64 `yaml:"input_per_1k"`
|
||||
OutputPer1K float64 `yaml:"output_per_1k"`
|
||||
CachedInputPer1K float64 `yaml:"cached_input_per_1k"`
|
||||
CacheReadPer1K float64 `yaml:"cache_read_per_1k"`
|
||||
CacheCreationPer1K float64 `yaml:"cache_creation_per_1k"`
|
||||
}
|
||||
|
||||
const (
|
||||
// ReloadInterval is the mtime-poll cadence for the background reloader.
|
||||
ReloadInterval = 30 * time.Second
|
||||
|
||||
// errorBackoff bounds how often the loader logs a repeated parse error.
|
||||
errorBackoff = 5 * time.Minute
|
||||
)
|
||||
|
||||
var basenameRegex = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`)
|
||||
|
||||
// Loader is a confined, hot-reloadable pricing table reader. Construction
|
||||
// must succeed against the target file; subsequent reload failures keep the
|
||||
// previously-loaded table so callers never observe a blank price list.
|
||||
type Loader struct {
|
||||
baseDir string
|
||||
fullPath string
|
||||
pluginID string
|
||||
table atomic.Pointer[Table]
|
||||
mtime atomic.Int64
|
||||
failures metric.Int64Counter
|
||||
interval time.Duration
|
||||
}
|
||||
|
||||
// NewLoader returns a pricing loader that overlays an optional file-based
|
||||
// table on top of the embedded defaults. Missing override file, baseDir, or
|
||||
// relPath is not an error: the loader keeps the embedded defaults so cost
|
||||
// metadata is still emitted for known models.
|
||||
//
|
||||
// Errors:
|
||||
// - bad basename, traversal segment, or absolute relPath are rejected so a
|
||||
// misconfigured target surfaces immediately.
|
||||
// - permission errors and YAML parse errors keep the defaults but log a
|
||||
// warning; cost annotation does not silently break.
|
||||
//
|
||||
// failures is optional; pass nil in tests that do not care about
|
||||
// reload-failure telemetry.
|
||||
func NewLoader(baseDir, relPath, pluginID string, failures metric.Int64Counter) (*Loader, error) {
|
||||
defaults := DefaultTable()
|
||||
l := &Loader{
|
||||
baseDir: baseDir,
|
||||
pluginID: pluginID,
|
||||
failures: failures,
|
||||
}
|
||||
l.table.Store(defaults)
|
||||
|
||||
if strings.TrimSpace(baseDir) == "" || strings.TrimSpace(relPath) == "" {
|
||||
return l, nil
|
||||
}
|
||||
|
||||
full, err := resolveMiddlewareDataPath(baseDir, relPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
l.fullPath = full
|
||||
|
||||
overlay, mtime, err := loadPricing(full)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// Override file is optional. Defaults already stored.
|
||||
return l, nil
|
||||
}
|
||||
// Symlink rejection, oversize file, parse failure, permission errors
|
||||
// — surface so a misconfigured operator sees the problem instead of
|
||||
// silently running with stale defaults.
|
||||
return nil, fmt.Errorf("load pricing %s: %w", full, err)
|
||||
}
|
||||
l.table.Store(mergeOver(defaults, overlay))
|
||||
l.mtime.Store(mtime.UnixNano())
|
||||
return l, nil
|
||||
}
|
||||
|
||||
// Get returns the current pricing table. The returned pointer is immutable;
|
||||
// callers must not mutate its contents.
|
||||
func (l *Loader) Get() *Table {
|
||||
if l == nil {
|
||||
return nil
|
||||
}
|
||||
return l.table.Load()
|
||||
}
|
||||
|
||||
// WatchesFile reports whether this loader is bound to an override file on
|
||||
// disk. False for defaults-only loaders (no operator override given).
|
||||
// Callers use this to decide whether to spawn the mtime-poll goroutine.
|
||||
func (l *Loader) WatchesFile() bool {
|
||||
if l == nil {
|
||||
return false
|
||||
}
|
||||
return l.fullPath != ""
|
||||
}
|
||||
|
||||
// SetReloadInterval overrides the mtime-poll cadence used by Reload. Calls
|
||||
// after Reload has started have no effect on the running loop. Intended for
|
||||
// tests; production code uses the default ReloadInterval.
|
||||
func (l *Loader) SetReloadInterval(d time.Duration) {
|
||||
if l == nil || d <= 0 {
|
||||
return
|
||||
}
|
||||
l.interval = d
|
||||
}
|
||||
|
||||
// Reload runs a polling loop that checks the pricing file mtime every
|
||||
// ReloadInterval (or the value passed to SetReloadInterval). Returns when
|
||||
// ctx is cancelled.
|
||||
func (l *Loader) Reload(ctx context.Context) {
|
||||
if l == nil {
|
||||
return
|
||||
}
|
||||
interval := l.interval
|
||||
if interval <= 0 {
|
||||
interval = ReloadInterval
|
||||
}
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
|
||||
var lastErrAt time.Time
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
if err := l.reload(); err != nil {
|
||||
if l.failures != nil {
|
||||
l.failures.Add(ctx, 1, metric.WithAttributes(
|
||||
attribute.String("plugin", l.pluginID),
|
||||
))
|
||||
}
|
||||
now := time.Now()
|
||||
if now.Sub(lastErrAt) >= errorBackoff {
|
||||
log.Warnf("llmobs: pricing reload failed for %s: %v", l.fullPath, err)
|
||||
lastErrAt = now
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// reload performs a single-shot mtime check and reload. The reloaded
|
||||
// override file is merged on top of the embedded defaults; missing override
|
||||
// (e.g. operator deleted the file) is not an error and reverts to defaults.
|
||||
func (l *Loader) reload() error {
|
||||
if l.fullPath == "" {
|
||||
// Defaults-only loader; nothing on disk to reload.
|
||||
return nil
|
||||
}
|
||||
mtime, err := statMtime(l.fullPath)
|
||||
if err != nil {
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
// File was removed since startup. Drop back to defaults and
|
||||
// reset mtime so a future re-creation triggers a reload.
|
||||
l.table.Store(DefaultTable())
|
||||
l.mtime.Store(0)
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if mtime.UnixNano() == l.mtime.Load() {
|
||||
return nil
|
||||
}
|
||||
|
||||
overlay, newMtime, err := loadPricing(l.fullPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
l.table.Store(mergeOver(DefaultTable(), overlay))
|
||||
l.mtime.Store(newMtime.UnixNano())
|
||||
return nil
|
||||
}
|
||||
|
||||
// resolveMiddlewareDataPath validates relPath is a safe basename and resolves
|
||||
// it under baseDir. An additional cleaned-prefix check guards against
|
||||
// CVE-style edge cases where Join is used with trailing path segments.
|
||||
func resolveMiddlewareDataPath(baseDir, relPath string) (string, error) {
|
||||
if strings.TrimSpace(baseDir) == "" {
|
||||
return "", errors.New("middleware-data-dir is not configured")
|
||||
}
|
||||
if relPath == "" {
|
||||
return "", errors.New("pricing path is empty")
|
||||
}
|
||||
if !basenameRegex.MatchString(relPath) {
|
||||
return "", fmt.Errorf("pricing path %q is not a safe basename", relPath)
|
||||
}
|
||||
if filepath.IsAbs(relPath) {
|
||||
return "", fmt.Errorf("pricing path %q must be a basename, not absolute", relPath)
|
||||
}
|
||||
|
||||
cleanBase, err := filepath.Abs(filepath.Clean(baseDir))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolve middleware-data-dir: %w", err)
|
||||
}
|
||||
full := filepath.Join(cleanBase, relPath)
|
||||
cleanedFull := filepath.Clean(full)
|
||||
if !strings.HasPrefix(cleanedFull, cleanBase+string(filepath.Separator)) && cleanedFull != cleanBase {
|
||||
return "", fmt.Errorf("pricing path %q escapes middleware-data-dir", relPath)
|
||||
}
|
||||
return cleanedFull, nil
|
||||
}
|
||||
|
||||
func parsePricingBytes(data []byte) (*Table, error) {
|
||||
dec := yaml.NewDecoder(bytes.NewReader(data))
|
||||
dec.KnownFields(true)
|
||||
|
||||
var raw pricingFile
|
||||
if err := dec.Decode(&raw); err != nil && !errors.Is(err, io.EOF) {
|
||||
return nil, fmt.Errorf("decode pricing yaml: %w", err)
|
||||
}
|
||||
|
||||
out := make(map[string]map[string]Entry, len(raw))
|
||||
for provider, models := range raw {
|
||||
inner := make(map[string]Entry, len(models))
|
||||
for model, entry := range models {
|
||||
for field, v := range map[string]float64{
|
||||
"input_per_1k": entry.InputPer1K,
|
||||
"output_per_1k": entry.OutputPer1K,
|
||||
"cached_input_per_1k": entry.CachedInputPer1K,
|
||||
"cache_read_per_1k": entry.CacheReadPer1K,
|
||||
"cache_creation_per_1k": entry.CacheCreationPer1K,
|
||||
} {
|
||||
if v < 0 || math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return nil, fmt.Errorf("pricing %s/%s: %s must be a finite, non-negative rate, got %v", provider, model, field, v)
|
||||
}
|
||||
}
|
||||
inner[model] = Entry{
|
||||
InputPer1K: entry.InputPer1K,
|
||||
OutputPer1K: entry.OutputPer1K,
|
||||
CachedInputPer1K: entry.CachedInputPer1K,
|
||||
CacheReadPer1K: entry.CacheReadPer1K,
|
||||
CacheCreationPer1K: entry.CacheCreationPer1K,
|
||||
}
|
||||
}
|
||||
out[provider] = inner
|
||||
}
|
||||
return &Table{entries: out}, nil
|
||||
}
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
//go:build !unix
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// loadPricing is unavailable on non-Unix platforms because O_NOFOLLOW and
|
||||
// fstat-from-FD are required to honour the spec's symlink-safety rules. The
|
||||
// proxy is only deployed on Linux today; a Windows port would need an
|
||||
// equivalent path-as-handle implementation.
|
||||
func loadPricing(path string) (*Table, time.Time, error) {
|
||||
return nil, time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path)
|
||||
}
|
||||
|
||||
func statMtime(path string) (time.Time, error) {
|
||||
return time.Time{}, fmt.Errorf("llmobs pricing loader is not supported on this platform: %s", path)
|
||||
}
|
||||
@@ -1,47 +1,13 @@
|
||||
//go:build unix
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"math"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func copyFixture(t *testing.T, src, dst string) {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile(src)
|
||||
require.NoError(t, err, "read source fixture")
|
||||
require.NoError(t, os.WriteFile(dst, data, 0o600), "write target fixture")
|
||||
}
|
||||
|
||||
func TestNewLoader_HappyPath(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err, "NewLoader must succeed with a valid fixture")
|
||||
table := l.Get()
|
||||
require.NotNil(t, table, "table populated after load")
|
||||
|
||||
cost, ok := table.Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "known provider/model resolves")
|
||||
assert.InDelta(t, 0.00075, cost, 1e-9, "cost = 0.00015 + 0.0006 per 1k tokens")
|
||||
|
||||
cost, ok = table.Cost("openai", "gpt-4o", 2000, 1000, 0, 0)
|
||||
require.True(t, ok, "second known model resolves")
|
||||
assert.InDelta(t, 0.015, cost, 1e-9, "cost for gpt-4o: 2*0.0025 + 1*0.01")
|
||||
|
||||
cost, ok = table.Cost("anthropic", "claude-sonnet-4-5", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "anthropic model resolves")
|
||||
assert.InDelta(t, 0.018, cost, 1e-9, "cost for claude-sonnet-4-5: 0.003 + 0.015")
|
||||
}
|
||||
|
||||
// TestCost_OpenAICachedSubsetDiscount proves OpenAI's cached input
|
||||
// tokens are billed at the configured cached_input_per_1k rate while
|
||||
// the non-cached remainder of input_tokens is billed at the regular
|
||||
@@ -65,11 +31,9 @@ func TestCost_OpenAICachedSubsetDiscount(t *testing.T) {
|
||||
"cached subset must bill at the discount rate; non-cached remainder at regular rate")
|
||||
}
|
||||
|
||||
// TestCost_OpenAICachedFallsBackToInputRate covers the operator
|
||||
// opt-in contract: when CachedInputPer1K is unset (zero), cached
|
||||
// tokens bill at the regular input rate. This matches today's
|
||||
// behaviour (cached counts weren't extracted at all so they
|
||||
// implicitly billed at the input rate via prompt_tokens).
|
||||
// TestCost_OpenAICachedFallsBackToInputRate covers the fallback
|
||||
// contract: when CachedInputPer1K is unset (zero), cached tokens bill
|
||||
// at the regular input rate.
|
||||
func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) {
|
||||
tbl := &Table{entries: map[string]map[string]Entry{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}},
|
||||
@@ -78,7 +42,7 @@ func TestCost_OpenAICachedFallsBackToInputRate(t *testing.T) {
|
||||
require.True(t, ok)
|
||||
want := 0.0025 + (500.0/1000.0)*0.01
|
||||
assert.InDelta(t, want, cost, 1e-12,
|
||||
"absent cached_input_per_1k rate must fall back to input_per_1k — same as pre-feature behaviour")
|
||||
"absent cached_input_per_1k rate must fall back to input_per_1k")
|
||||
}
|
||||
|
||||
// TestCost_OpenAIClampsCachedToInputCount is the defensive guard
|
||||
@@ -100,10 +64,7 @@ func TestCost_OpenAIClampsCachedToInputCount(t *testing.T) {
|
||||
// TestCost_AnthropicCacheReadAndCreationAreAdditive proves the
|
||||
// Anthropic shape: cache_read and cache_creation tokens are
|
||||
// ADDITIVE to input_tokens (not subset), each billed at its own
|
||||
// configured rate. The two rates pull in opposite directions —
|
||||
// cache_read is the cheaper read-from-cache rate (≈0.1× input),
|
||||
// cache_creation is the more expensive write-to-cache rate
|
||||
// (≈1.25× input).
|
||||
// configured rate.
|
||||
func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) {
|
||||
tbl := &Table{entries: map[string]map[string]Entry{
|
||||
"anthropic": {"claude-sonnet": {
|
||||
@@ -125,11 +86,9 @@ func TestCost_AnthropicCacheReadAndCreationAreAdditive(t *testing.T) {
|
||||
"each Anthropic input bucket must bill at its own configured rate")
|
||||
}
|
||||
|
||||
// TestCost_AnthropicCacheRatesFallBackToInput covers the no-opt-in
|
||||
// TestCost_AnthropicCacheRatesFallBackToInput covers the no-rate
|
||||
// path: when neither CacheReadPer1K nor CacheCreationPer1K is set,
|
||||
// cache tokens bill at the regular input rate. This is more
|
||||
// accurate than today's behaviour (cache tokens ignored entirely)
|
||||
// without requiring operators to opt in via YAML.
|
||||
// cache tokens bill at the regular input rate.
|
||||
func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) {
|
||||
tbl := &Table{entries: map[string]map[string]Entry{
|
||||
"anthropic": {"claude-sonnet": {InputPer1K: 0.003, OutputPer1K: 0.015}},
|
||||
@@ -139,259 +98,39 @@ func TestCost_AnthropicCacheRatesFallBackToInput(t *testing.T) {
|
||||
// Without overrides: every input bucket at input_per_1k.
|
||||
want := ((256.0+768.0+512.0)/1000.0)*0.003 + (200.0/1000.0)*0.015
|
||||
assert.InDelta(t, want, cost, 1e-12,
|
||||
"absent cache rates must fall back to input_per_1k — Anthropic cache tokens were ignored before this change, billing at input rate is more accurate as a default")
|
||||
"absent cache rates must fall back to input_per_1k")
|
||||
}
|
||||
|
||||
func TestNewLoader_UnknownModel(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
// TestEntryCosts_SurfaceSelectsFormula pins that the formula branches on
|
||||
// the SURFACE, not on which table the entry came from: the same entry
|
||||
// bills a subset carve-out on "openai", additive buckets on
|
||||
// "anthropic"/"bedrock", and ignores cache counts everywhere else. This
|
||||
// is what keeps per-provider-record entries (looked up by record id)
|
||||
// mathematically identical to defaults-table entries.
|
||||
func TestEntryCosts_SurfaceSelectsFormula(t *testing.T) {
|
||||
e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01, CachedInputPer1K: 0.001, CacheReadPer1K: 0.0002, CacheCreationPer1K: 0.0025}
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
openai := EntryCosts(e, "openai", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, (600.0/1000.0)*0.002+(400.0/1000.0)*0.001, openai.TotalUSD, 1e-12,
|
||||
"openai: cached is a subset, cacheCreation ignored")
|
||||
|
||||
_, ok := l.Get().Cost("openai", "fantasy-model", 10, 10, 0, 0)
|
||||
assert.False(t, ok, "unknown model returns ok=false")
|
||||
anthropic := EntryCosts(e, "anthropic", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, 0.002+(400.0/1000.0)*0.0002+(300.0/1000.0)*0.0025, anthropic.TotalUSD, 1e-12,
|
||||
"anthropic: cache buckets are additive")
|
||||
|
||||
_, ok = l.Get().Cost("cohere", "anything", 10, 10, 0, 0)
|
||||
assert.False(t, ok, "unknown provider returns ok=false")
|
||||
bedrock := EntryCosts(e, "bedrock", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, anthropic.TotalUSD, bedrock.TotalUSD, 1e-12, "bedrock shares the anthropic formula")
|
||||
|
||||
other := EntryCosts(e, "gemini", 1000, 0, 400, 300)
|
||||
assert.InDelta(t, 0.002, other.TotalUSD, 1e-12, "unknown surface: cache counts ignored")
|
||||
}
|
||||
|
||||
func TestNewLoader_InvalidYAMLRejected(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(base, "pricing.yaml"), []byte("\t- this is not: valid: yaml: :["), 0o600))
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "invalid YAML must surface as construction error")
|
||||
}
|
||||
|
||||
func TestLoader_ReloadKeepsPreviousOnParseError(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, l.Get(), "initial table populated")
|
||||
|
||||
// Overwrite with content that violates the strict schema (extra field)
|
||||
// plus a bumped mtime to trigger reload.
|
||||
require.NoError(t, os.WriteFile(target, []byte("openai:\n gpt-4o:\n input_per_1k: 1.0\n output_per_1k: 2.0\n bogus_field: nope\n"), 0o600))
|
||||
future := time.Now().Add(time.Hour)
|
||||
require.NoError(t, os.Chtimes(target, future, future))
|
||||
|
||||
err = l.reload()
|
||||
require.Error(t, err, "parse error surfaced by reload()")
|
||||
|
||||
cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "previous table still available after parse failure")
|
||||
assert.InDelta(t, 0.00075, cost, 1e-9, "previous cost preserved")
|
||||
}
|
||||
|
||||
func TestLoader_ReloadNoChangeIsNoOp(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
ptrBefore := l.Get()
|
||||
|
||||
require.NoError(t, l.reload(), "no-change reload must not error")
|
||||
ptrAfter := l.Get()
|
||||
assert.Same(t, ptrBefore, ptrAfter, "table pointer unchanged when mtime unchanged")
|
||||
}
|
||||
|
||||
func TestLoader_ReloadDetectsChange(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n")
|
||||
require.NoError(t, os.WriteFile(target, updated, 0o600))
|
||||
future := time.Now().Add(time.Hour)
|
||||
require.NoError(t, os.Chtimes(target, future, future))
|
||||
|
||||
require.NoError(t, l.reload(), "reload must succeed on valid new content")
|
||||
|
||||
cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "updated model still present")
|
||||
assert.InDelta(t, 3.0, cost, 0.0001, "new prices are applied: 1 + 2 per 1k")
|
||||
}
|
||||
|
||||
// TestLoader_ReloadGoroutinePicksUpChanges proves the background goroutine
|
||||
// started via Reload actually swaps the pricing table when the file changes
|
||||
// on disk. Without that goroutine running, pricing edits would never reach
|
||||
// requests until a proxy restart.
|
||||
func TestLoader_ReloadGoroutinePicksUpChanges(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
l.SetReloadInterval(20 * time.Millisecond)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
l.Reload(ctx)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Before any rewrite, the loader holds the fixture's prices.
|
||||
costBefore, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "fixture model must resolve initially")
|
||||
assert.InDelta(t, 0.00075, costBefore, 1e-9, "fixture prices apply before rewrite")
|
||||
|
||||
updated := []byte("openai:\n gpt-4o-mini:\n input_per_1k: 1.00\n output_per_1k: 2.00\n")
|
||||
require.NoError(t, os.WriteFile(target, updated, 0o600))
|
||||
future := time.Now().Add(time.Hour)
|
||||
require.NoError(t, os.Chtimes(target, future, future))
|
||||
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for {
|
||||
cost, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
if ok && cost > 2.5 {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("background reloader did not pick up rewrite within deadline")
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("Reload loop did not exit after cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoader_ReloadBackgroundLoopCancellation(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
l, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
l.Reload(ctx)
|
||||
close(done)
|
||||
}()
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("Reload loop did not exit on context cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewLoader_PathValidation(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
relPath string
|
||||
}{
|
||||
{"traversal", "../../etc/passwd"},
|
||||
{"absolute", "/etc/passwd"},
|
||||
{"slash in basename", "sub/pricing.yaml"},
|
||||
{"control chars", "pricing\x00.yaml"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewLoader(base, tc.relPath, "llm_observability", nil)
|
||||
require.Error(t, err, "NewLoader must reject %q", tc.relPath)
|
||||
})
|
||||
}
|
||||
|
||||
// Empty relPath is no longer a validation error: the loader treats it
|
||||
// as "no override file, defaults only" so cost metadata is still
|
||||
// emitted for the embedded models out of the box.
|
||||
t.Run("empty falls back to defaults", func(t *testing.T) {
|
||||
l, err := NewLoader(base, "", "llm_observability", nil)
|
||||
require.NoError(t, err, "empty relPath should yield a defaults-only loader")
|
||||
require.NotNil(t, l, "loader must be returned")
|
||||
require.False(t, l.WatchesFile(), "no file watching when no override is given")
|
||||
_, ok := l.Get().Cost("openai", "gpt-4o-mini", 1000, 1000, 0, 0)
|
||||
assert.True(t, ok, "embedded defaults should still resolve gpt-4o-mini")
|
||||
})
|
||||
}
|
||||
|
||||
// TestNewLoader_PathValidation_Extended covers the remaining attack shapes
|
||||
// called out in C2: dot references, embedded traversal segments, and a
|
||||
// newline in the basename. The basename regex must reject each one even
|
||||
// though filepath.Clean would otherwise collapse them.
|
||||
func TestNewLoader_PathValidation_Extended(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing.yaml"))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
relPath string
|
||||
}{
|
||||
{"dot", "."},
|
||||
{"dotdot", ".."},
|
||||
{"relative traversal", "../pricing.yaml"},
|
||||
{"embedded slash", "pri/cing.yaml"},
|
||||
{"newline", "pricing\n.yaml"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewLoader(base, tc.relPath, "llm_observability", nil)
|
||||
require.Error(t, err, "NewLoader must reject %q", tc.relPath)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewLoader_ValidBasenameLoads proves the allowlist is exclusive: a
|
||||
// basename containing only safe characters under baseDir loads. Without this
|
||||
// a regression that over-tightened the regex would silently break valid
|
||||
// deployments.
|
||||
func TestNewLoader_ValidBasenameLoads(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), filepath.Join(base, "pricing-v2_prod.yaml"))
|
||||
|
||||
l, err := NewLoader(base, "pricing-v2_prod.yaml", "llm_observability", nil)
|
||||
require.NoError(t, err, "basename with _, -, . must load")
|
||||
require.NotNil(t, l.Get(), "table populated")
|
||||
}
|
||||
|
||||
// TestNewLoader_SymlinkOutsideBaseDirRejected constructs a symlink under
|
||||
// baseDir that points to a file outside it. O_NOFOLLOW must refuse to open
|
||||
// the symlink even though the symlink path itself is a valid basename under
|
||||
// baseDir.
|
||||
func TestNewLoader_SymlinkOutsideBaseDirRejected(t *testing.T) {
|
||||
outside := t.TempDir()
|
||||
target := filepath.Join(outside, "evil.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), target)
|
||||
|
||||
base := t.TempDir()
|
||||
link := filepath.Join(base, "pricing.yaml")
|
||||
require.NoError(t, os.Symlink(target, link), "symlink setup")
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "O_NOFOLLOW must reject symlink even when it points outside baseDir")
|
||||
}
|
||||
|
||||
func TestNewLoader_SymlinkRejected(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
concrete := filepath.Join(base, "real.yaml")
|
||||
copyFixture(t, filepath.Join("..", "fixtures", "pricing.yaml"), concrete)
|
||||
|
||||
link := filepath.Join(base, "pricing.yaml")
|
||||
require.NoError(t, os.Symlink(concrete, link), "symlink setup")
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "O_NOFOLLOW must reject symlinked targets")
|
||||
// TestEntryCosts_ClampsNegativeTokens: malformed upstream counts must
|
||||
// never produce a negative cost.
|
||||
func TestEntryCosts_ClampsNegativeTokens(t *testing.T) {
|
||||
e := Entry{InputPer1K: 0.002, OutputPer1K: 0.01}
|
||||
c := EntryCosts(e, "openai", -50, -10, -5, -3)
|
||||
assert.Zero(t, c.TotalUSD, "all-negative counts clamp to zero cost")
|
||||
}
|
||||
|
||||
func TestTableCost_NilSafe(t *testing.T) {
|
||||
@@ -402,31 +141,37 @@ func TestTableCost_NilSafe(t *testing.T) {
|
||||
assert.False(t, t1.Has("x", "y"), "nil table has nothing")
|
||||
}
|
||||
|
||||
func TestLoaderGet_NilSafe(t *testing.T) {
|
||||
var l *Loader
|
||||
assert.Nil(t, l.Get(), "nil loader returns nil table")
|
||||
}
|
||||
|
||||
// TestNewLoader_RejectsOversizedFile_FixesM4 proves the loader bounds reads
|
||||
// at maxPricingBytes so a hostile file cannot exhaust process memory.
|
||||
func TestNewLoader_RejectsOversizedFile_FixesM4(t *testing.T) {
|
||||
base := t.TempDir()
|
||||
target := filepath.Join(base, "pricing.yaml")
|
||||
|
||||
// Build a YAML payload larger than the cap. We pad with valid YAML
|
||||
// comments so a partial read would still fail the size check rather
|
||||
// than the parser.
|
||||
header := "openai:\n"
|
||||
bigComment := make([]byte, maxPricingBytes+1024)
|
||||
for i := range bigComment {
|
||||
bigComment[i] = ' '
|
||||
func TestNewTable_ValidatesRates(t *testing.T) {
|
||||
good := map[string]map[string]EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125}},
|
||||
}
|
||||
bigComment[0] = '#'
|
||||
bigComment[len(bigComment)-1] = '\n'
|
||||
payload := append([]byte(header), bigComment...)
|
||||
require.NoError(t, os.WriteFile(target, payload, 0o600))
|
||||
tbl, err := NewTable(good)
|
||||
require.NoError(t, err)
|
||||
cost, ok := tbl.Cost("openai", "gpt-4o", 1000, 1000, 0, 0)
|
||||
require.True(t, ok, "entry survives the wire conversion")
|
||||
assert.InDelta(t, 0.0125, cost, 1e-9)
|
||||
|
||||
_, err := NewLoader(base, "pricing.yaml", "llm_observability", nil)
|
||||
require.Error(t, err, "oversized pricing file must be rejected")
|
||||
assert.Contains(t, err.Error(), "exceeds", "rejection must reference the byte cap")
|
||||
_, ok = tbl.Cost("openai", "unknown-model", 1, 1, 0, 0)
|
||||
assert.False(t, ok, "unknown model misses")
|
||||
|
||||
for name, bad := range map[string]EntryJSON{
|
||||
"negative input": {InputPer1K: -1, OutputPer1K: 0.01},
|
||||
"NaN output": {InputPer1K: 0.01, OutputPer1K: math.NaN()},
|
||||
"Inf cache read": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: math.Inf(1)},
|
||||
"negative cached": {InputPer1K: 0.01, OutputPer1K: 0.01, CachedInputPer1K: -0.001},
|
||||
} {
|
||||
_, err := NewTable(map[string]map[string]EntryJSON{"openai": {"m": bad}})
|
||||
assert.Error(t, err, "case %q must be rejected so a corrupt config fails the chain build instead of mispricing", name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewTable_NilAndEmpty(t *testing.T) {
|
||||
tbl, err := NewTable(nil)
|
||||
require.NoError(t, err, "nil map builds an empty (never-matching) table")
|
||||
_, ok := tbl.Cost("openai", "gpt-4o", 1, 1, 0, 0)
|
||||
assert.False(t, ok, "empty table prices nothing")
|
||||
|
||||
entries, err := NewEntries(nil)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, entries, "nil in, empty (never-matching) map out for the per-record map")
|
||||
}
|
||||
|
||||
@@ -1,68 +0,0 @@
|
||||
//go:build unix
|
||||
|
||||
package pricing
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// maxPricingBytes caps the size of the pricing YAML on read so a hostile or
|
||||
// runaway file cannot exhaust process memory during reload. 1 MiB is several
|
||||
// orders of magnitude larger than any reasonable pricing table.
|
||||
const maxPricingBytes int64 = 1 << 20
|
||||
|
||||
// loadPricing opens the file with O_NOFOLLOW, fstats the open descriptor,
|
||||
// and parses from that same descriptor. Never re-opens by path so a
|
||||
// mid-read rename or symlink swap cannot substitute content. Bytes are
|
||||
// capped at maxPricingBytes so the loader cannot be coerced into reading an
|
||||
// unbounded file.
|
||||
func loadPricing(path string) (*Table, time.Time, error) {
|
||||
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
|
||||
if err != nil {
|
||||
return nil, time.Time{}, fmt.Errorf("open %s: %w", path, err)
|
||||
}
|
||||
defer func() {
|
||||
if cerr := f.Close(); cerr != nil {
|
||||
log.Debugf("close pricing file %s: %v", path, cerr)
|
||||
}
|
||||
}()
|
||||
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
return nil, time.Time{}, fmt.Errorf("fstat %s: %w", path, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return nil, time.Time{}, fmt.Errorf("pricing file %s is not a regular file", path)
|
||||
}
|
||||
|
||||
data, err := io.ReadAll(io.LimitReader(f, maxPricingBytes+1))
|
||||
if err != nil {
|
||||
return nil, time.Time{}, fmt.Errorf("read %s: %w", path, err)
|
||||
}
|
||||
if int64(len(data)) > maxPricingBytes {
|
||||
return nil, time.Time{}, fmt.Errorf("pricing file %s exceeds %d bytes", path, maxPricingBytes)
|
||||
}
|
||||
|
||||
table, err := parsePricingBytes(data)
|
||||
if err != nil {
|
||||
return nil, time.Time{}, err
|
||||
}
|
||||
return table, info.ModTime(), nil
|
||||
}
|
||||
|
||||
// statMtime returns the mtime of the file at path. It uses lstat semantics
|
||||
// via os.Lstat so a symlink swap is detected even though O_NOFOLLOW will
|
||||
// later reject the open.
|
||||
func statMtime(path string) (time.Time, error) {
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
return time.Time{}, fmt.Errorf("lstat %s: %w", path, err)
|
||||
}
|
||||
return info.ModTime(), nil
|
||||
}
|
||||
@@ -36,15 +36,13 @@ var defaultRegistry = middleware.NewRegistry()
|
||||
|
||||
// FactoryContext is the per-process bag that concrete factories may
|
||||
// consult during construction. It carries the proxy-lifetime context,
|
||||
// the data directory used for static config files (pricing tables,
|
||||
// allowlists), the OTel meter, and the proxy logger.
|
||||
// the OTel meter, and the proxy logger.
|
||||
//
|
||||
// Configure must be called once at boot before any chain build calls
|
||||
// Resolve. Calling it twice overwrites the prior value; tests may rely
|
||||
// on this to reset state.
|
||||
type FactoryContext struct {
|
||||
Context context.Context
|
||||
DataDir string
|
||||
Meter metric.Meter
|
||||
Logger *log.Logger
|
||||
MgmtClient MgmtClient
|
||||
@@ -58,12 +56,11 @@ var (
|
||||
// Configure stores the per-process FactoryContext. Concrete factories
|
||||
// reach for it via Context(). mgmt may be nil on tests / standalone
|
||||
// builds with no management server; consumers must guard.
|
||||
func Configure(ctx context.Context, dataDir string, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) {
|
||||
func Configure(ctx context.Context, meter metric.Meter, logger *log.Logger, mgmt MgmtClient) {
|
||||
ctxMu.Lock()
|
||||
defer ctxMu.Unlock()
|
||||
ctxStore = FactoryContext{
|
||||
Context: ctx,
|
||||
DataDir: dataDir,
|
||||
Meter: meter,
|
||||
Logger: logger,
|
||||
MgmtClient: mgmt,
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
mgmtpricing "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/cost_meter"
|
||||
@@ -19,17 +20,20 @@ import (
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_response_parser"
|
||||
)
|
||||
|
||||
// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the embedded default pricing
|
||||
// table and asserts exact USD amounts hardcoded from the vendors' published prices, including the cache split.
|
||||
// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the REAL default pricing
|
||||
// table management ships (mgmtpricing.DefaultTable, catalog-derived) and asserts exact USD amounts hardcoded from
|
||||
// the vendors' published prices, including the cache split. This is the cross-stack pricing contract test: the
|
||||
// management-side Entry JSON must decode into the proxy-side table and produce these exact costs.
|
||||
func TestCostCalculation_ProviderMatrix(t *testing.T) {
|
||||
// Empty data dir → embedded defaults, like a proxy with no pricing override.
|
||||
builtin.Configure(context.Background(), t.TempDir(), nil, nil, nil)
|
||||
builtin.Configure(context.Background(), nil, nil, nil)
|
||||
|
||||
reqMW, err := llm_request_parser.Factory{}.New(nil)
|
||||
require.NoError(t, err, "build llm_request_parser")
|
||||
respMW, err := llm_response_parser.Factory{}.New(nil)
|
||||
require.NoError(t, err, "build llm_response_parser")
|
||||
costMW, err := cost_meter.Factory{}.New(nil)
|
||||
costCfgJSON, err := json.Marshal(map[string]any{"pricing": map[string]any{"defaults": mgmtpricing.DefaultTable()}})
|
||||
require.NoError(t, err, "marshal management default table into cost_meter config")
|
||||
costMW, err := cost_meter.Factory{}.New(costCfgJSON)
|
||||
require.NoError(t, err, "build cost_meter")
|
||||
t.Cleanup(func() { _ = costMW.Close() })
|
||||
|
||||
|
||||
@@ -2,7 +2,6 @@ package cost_meter
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
@@ -11,16 +10,27 @@ import (
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
)
|
||||
|
||||
// defaultPricingFilename is the basename probed inside the proxy data
|
||||
// directory when no override is configured.
|
||||
const defaultPricingFilename = "pricing.yaml"
|
||||
|
||||
// Config is the on-wire configuration for the middleware.
|
||||
// Config is the on-wire configuration for the middleware, synthesized by
|
||||
// management (buildCostMeterConfigJSON). The proxy has no embedded price
|
||||
// list: this payload is the only pricing source, and updates arrive as
|
||||
// ordinary mapping pushes that rebuild the chain (and with it this
|
||||
// middleware instance) — no per-request fetches, no reload loops.
|
||||
type Config struct {
|
||||
// PricingPath optionally overrides the basename of the pricing
|
||||
// file probed inside the proxy data directory. When empty the
|
||||
// loader falls back to "pricing.yaml".
|
||||
PricingPath string `json:"pricing_path"`
|
||||
Pricing *PricingConfig `json:"pricing"`
|
||||
}
|
||||
|
||||
// PricingConfig carries the full pricing table:
|
||||
// - Defaults: parser surface ("openai"/"anthropic"/"bedrock") ->
|
||||
// normalized model id -> rates, matched against llm.provider +
|
||||
// llm.model.
|
||||
// - Providers: provider record id -> normalized model id -> rates,
|
||||
// matched against the llm.resolved_provider_id metadata llm_router
|
||||
// stamps. Entries arrive fully materialized (management folds default
|
||||
// cache rates in at synth time), so lookup order is simply
|
||||
// per-record first, defaults second.
|
||||
type PricingConfig struct {
|
||||
Defaults map[string]map[string]pricing.EntryJSON `json:"defaults"`
|
||||
Providers map[string]map[string]pricing.EntryJSON `json:"providers"`
|
||||
}
|
||||
|
||||
// Factory builds cost_meter instances from raw config bytes.
|
||||
@@ -29,45 +39,45 @@ type Factory struct{}
|
||||
// ID returns the registry identifier.
|
||||
func (Factory) ID() string { return ID }
|
||||
|
||||
// New constructs a middleware instance. Empty, null, and {} configs
|
||||
// are accepted; non-empty rawConfig that fails to unmarshal is
|
||||
// rejected so misconfigurations surface at chain build time. The
|
||||
// pricing loader is built once per instance and reused across
|
||||
// invocations.
|
||||
// New constructs a middleware instance. Empty, null, and {} configs are
|
||||
// accepted for backward compatibility with a management server that
|
||||
// predates config-delivered pricing — the instance then skips every cost
|
||||
// computation (unknown_model) and a warning is logged once at build time.
|
||||
// Non-empty rawConfig that fails to unmarshal, or a table carrying a
|
||||
// non-finite / negative rate, is rejected so misconfigurations surface at
|
||||
// chain build time.
|
||||
func (Factory) New(rawConfig []byte) (middleware.Middleware, error) {
|
||||
cfg, err := decodeConfig(rawConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
fctx := builtin.Context()
|
||||
pricingPath := cfg.PricingPath
|
||||
if pricingPath == "" {
|
||||
pricingPath = defaultPricingFilename
|
||||
if cfg.Pricing == nil {
|
||||
if logger := builtin.Context().Logger; logger != nil {
|
||||
logger.Warnf("cost_meter: no pricing table in middleware config; management predates config-delivered pricing — every request will record cost.skipped=unknown_model ($0)")
|
||||
}
|
||||
return newMiddleware(mustEmptyTable(), nil), nil
|
||||
}
|
||||
|
||||
loader, err := pricing.NewLoader(fctx.DataDir, pricingPath, ID, nil)
|
||||
defaults, err := pricing.NewTable(cfg.Pricing.Defaults)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("init pricing loader: %w", err)
|
||||
return nil, fmt.Errorf("cost_meter pricing defaults: %w", err)
|
||||
}
|
||||
|
||||
cancel := startReloader(fctx.Context, loader)
|
||||
|
||||
return newMiddleware(loader, cancel), nil
|
||||
perRecord, err := pricing.NewEntries(cfg.Pricing.Providers)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cost_meter per-provider pricing: %w", err)
|
||||
}
|
||||
return newMiddleware(defaults, perRecord), nil
|
||||
}
|
||||
|
||||
// startReloader binds the loader's mtime-poll goroutine to a context
|
||||
// derived from the proxy-lifetime context and returns its cancel func so
|
||||
// the owning middleware can stop the goroutine on teardown. Returns nil
|
||||
// when there's nothing to watch (nil context or defaults-only loader), in
|
||||
// which case the middleware's Close is a no-op.
|
||||
func startReloader(ctx context.Context, loader *pricing.Loader) context.CancelFunc {
|
||||
if ctx == nil || !loader.WatchesFile() {
|
||||
return nil
|
||||
// mustEmptyTable returns a valid empty table. NewTable on a nil map cannot
|
||||
// fail; the panic guard documents that invariant.
|
||||
func mustEmptyTable() *pricing.Table {
|
||||
t, err := pricing.NewTable(nil)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("cost_meter: empty pricing table must build: %v", err))
|
||||
}
|
||||
cctx, cancel := context.WithCancel(ctx)
|
||||
go loader.Reload(cctx)
|
||||
return cancel
|
||||
return t
|
||||
}
|
||||
|
||||
// decodeConfig accepts empty, null, and {} configs, returning a
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
// Package cost_meter implements the SlotOnResponse middleware that
|
||||
// converts token-usage metadata emitted by llm_response_parser into a
|
||||
// per-request USD cost estimate. The middleware uses the shared pricing
|
||||
// loader so operator pricing overrides apply to the chain.
|
||||
// per-request USD cost estimate. Pricing arrives from management inside
|
||||
// the middleware config: a per-provider-record table (the operator's
|
||||
// stored prices, matched via llm.resolved_provider_id) consulted first,
|
||||
// then the surface-keyed defaults table.
|
||||
package cost_meter
|
||||
|
||||
import (
|
||||
@@ -17,7 +19,9 @@ import (
|
||||
const ID = "cost_meter"
|
||||
|
||||
// Version is the implementation version emitted via the spec merge.
|
||||
const Version = "1.0.0"
|
||||
// 1.1.0: pricing is config-delivered (defaults + per-provider-record
|
||||
// entries) instead of proxy-embedded.
|
||||
const Version = "1.1.0"
|
||||
|
||||
// Skip reasons emitted under KeyCostSkipped. The set is closed; the
|
||||
// dashboard surfaces these verbatim.
|
||||
@@ -42,19 +46,21 @@ var metadataKeys = []string{
|
||||
}
|
||||
|
||||
// Middleware computes a per-response cost estimate from the token
|
||||
// counts emitted upstream by llm_response_parser.
|
||||
// counts emitted upstream by llm_response_parser. Both tables are
|
||||
// immutable — a pricing change arrives as a mapping push that rebuilds
|
||||
// the chain with a fresh instance.
|
||||
type Middleware struct {
|
||||
loader *pricing.Loader
|
||||
// cancel stops this instance's pricing-reload goroutine. Non-nil only
|
||||
// when the loader watches an override file; Close calls it so a chain
|
||||
// rebuild doesn't leak a poll goroutine per retired instance.
|
||||
cancel context.CancelFunc
|
||||
// defaults is the surface-keyed table (llm.provider x llm.model).
|
||||
defaults *pricing.Table
|
||||
// perRecord is keyed by provider record id (llm.resolved_provider_id)
|
||||
// then normalized model id; entries arrive fully materialized from
|
||||
// management. Consulted before defaults. May be nil.
|
||||
perRecord map[string]map[string]pricing.Entry
|
||||
}
|
||||
|
||||
// newMiddleware constructs a Middleware bound to the given pricing loader.
|
||||
// cancel may be nil (defaults-only loader with no reloader to stop).
|
||||
func newMiddleware(loader *pricing.Loader, cancel context.CancelFunc) *Middleware {
|
||||
return &Middleware{loader: loader, cancel: cancel}
|
||||
// newMiddleware constructs a Middleware over the given pricing tables.
|
||||
func newMiddleware(defaults *pricing.Table, perRecord map[string]map[string]pricing.Entry) *Middleware {
|
||||
return &Middleware{defaults: defaults, perRecord: perRecord}
|
||||
}
|
||||
|
||||
// ID returns the registry identifier.
|
||||
@@ -79,16 +85,9 @@ func (m *Middleware) MetadataKeys() []string {
|
||||
// response.
|
||||
func (m *Middleware) MutationsSupported() bool { return false }
|
||||
|
||||
// Close stops this instance's pricing-reload goroutine, if any. Called by
|
||||
// the chain when a rebuild retires the instance, so the mtime-poll loop
|
||||
// doesn't outlive the chain it belonged to. Safe to call on a nil receiver
|
||||
// and on an instance with no reloader.
|
||||
func (m *Middleware) Close() error {
|
||||
if m != nil && m.cancel != nil {
|
||||
m.cancel()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
// Close releases resources owned by the middleware. Stateless — the
|
||||
// pricing tables are plain maps owned by this instance.
|
||||
func (m *Middleware) Close() error { return nil }
|
||||
|
||||
// Invoke reads provider, model, and token metadata, looks up pricing,
|
||||
// and emits either KeyCostUSDTotal or KeyCostSkipped. The decision is
|
||||
@@ -144,8 +143,7 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
|
||||
return out, nil
|
||||
}
|
||||
|
||||
table := m.loader.Get()
|
||||
costs, ok := table.Costs(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
costs, ok := m.lookupCosts(in.Metadata, provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
if !ok {
|
||||
out.Metadata = skip(skipUnknownModel)
|
||||
return out, nil
|
||||
@@ -164,6 +162,26 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// lookupCosts resolves the price for this request and computes the cost
|
||||
// split. Resolution order:
|
||||
//
|
||||
// 1. Per-provider-record entry: the operator's stored price for the
|
||||
// provider route that served the request, keyed by the
|
||||
// llm.resolved_provider_id metadata llm_router stamped on the allow
|
||||
// path. Absent metadata (e.g. no router in the chain) skips this tier.
|
||||
// 2. Surface defaults: the catalog-derived table keyed by llm.provider.
|
||||
//
|
||||
// The surface always selects the cache formula — a per-record entry for an
|
||||
// Anthropic route still bills its cache buckets additively.
|
||||
func (m *Middleware) lookupCosts(md []middleware.KV, surface, model string, inTokens, outTokens, cachedTokens, cacheCreationTokens int64) (pricing.Costs, bool) {
|
||||
if recordID := lookupKV(md, middleware.KeyLLMResolvedProviderID); recordID != "" {
|
||||
if entry, ok := m.perRecord[recordID][model]; ok {
|
||||
return pricing.EntryCosts(entry, surface, inTokens, outTokens, cachedTokens, cacheCreationTokens), true
|
||||
}
|
||||
}
|
||||
return m.defaults.Costs(surface, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
|
||||
}
|
||||
|
||||
// usd renders a cost as the fixed-precision string every cost.usd_* key
|
||||
// carries, so the per-bucket values and the aggregates round identically.
|
||||
//
|
||||
|
||||
@@ -3,39 +3,50 @@ package cost_meter
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/llm/pricing"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware"
|
||||
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
|
||||
)
|
||||
|
||||
const fixturePricing = `openai:
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
gpt-4o-mini:
|
||||
input_per_1k: 0.00015
|
||||
output_per_1k: 0.0006
|
||||
anthropic:
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
`
|
||||
|
||||
// configureBuiltin points the package-level FactoryContext at a tmp
|
||||
// directory containing the test pricing fixture. Returns the path so
|
||||
// callers can override files later if needed.
|
||||
func configureBuiltin(t *testing.T) string {
|
||||
// fixtureConfig mirrors what management's buildCostMeterConfigJSON ships:
|
||||
// a surface-keyed defaults table. Rates match the retired YAML fixture so
|
||||
// every cost assertion below is byte-identical to the pre-feature values.
|
||||
func fixtureConfig(t *testing.T) []byte {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricing), 0o600), "write pricing fixture")
|
||||
builtin.Configure(context.Background(), dir, nil, nil, nil)
|
||||
return dir
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {
|
||||
"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01},
|
||||
"gpt-4o-mini": {InputPer1K: 0.00015, OutputPer1K: 0.0006},
|
||||
},
|
||||
"anthropic": {
|
||||
"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015},
|
||||
},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
return raw
|
||||
}
|
||||
|
||||
// fixtureConfigWithCache adds the cache-rate fields.
|
||||
func fixtureConfigWithCache(t *testing.T) []byte {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {
|
||||
"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01, CachedInputPer1K: 0.00125},
|
||||
},
|
||||
"anthropic": {
|
||||
"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375},
|
||||
},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
return raw
|
||||
}
|
||||
|
||||
func metaValue(t *testing.T, kvs []middleware.KV, key string) (string, bool) {
|
||||
@@ -56,8 +67,7 @@ func buildMiddleware(t *testing.T, raw []byte) middleware.Middleware {
|
||||
}
|
||||
|
||||
func TestMiddleware_StaticSurface(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
assert.Equal(t, ID, mw.ID(), "ID must match the registered constant")
|
||||
assert.Equal(t, Version, mw.Version(), "Version must match the constant")
|
||||
@@ -79,8 +89,10 @@ func TestMiddleware_StaticSurface(t *testing.T) {
|
||||
assert.Equal(t, expected, keys, "metadata key allowlist must match the spec")
|
||||
}
|
||||
|
||||
// TestFactory_AcceptsEmptyAndJSONConfig: empty/null/{} configs are what an
|
||||
// old management (pre config-delivered pricing) sends — they must build a
|
||||
// working (all-skip) instance, never fail the chain.
|
||||
func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
cases := [][]byte{nil, {}, []byte("null"), []byte("{}"), []byte(" ")}
|
||||
for _, raw := range cases {
|
||||
mw, err := Factory{}.New(raw)
|
||||
@@ -90,15 +102,57 @@ func TestFactory_AcceptsEmptyAndJSONConfig(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFactory_RejectsMalformedConfig(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw, err := Factory{}.New([]byte("{not json"))
|
||||
require.Error(t, err, "malformed config must surface at construction")
|
||||
assert.Nil(t, mw, "no instance is returned on error")
|
||||
}
|
||||
|
||||
func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
// TestFactory_RejectsInvalidRates: a non-finite or negative rate anywhere
|
||||
// in the table fails the chain build (defense-in-depth behind management's
|
||||
// API validation) rather than silently mispricing.
|
||||
func TestFactory_RejectsInvalidRates(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: -0.0025, OutputPer1K: 0.01}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw, err := Factory{}.New(raw)
|
||||
require.Error(t, err, "negative rate must fail the build")
|
||||
assert.Nil(t, mw)
|
||||
|
||||
raw, err = json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-1": {"m": {InputPer1K: 0.01, OutputPer1K: 0.01, CacheReadPer1K: -1}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
_, err = Factory{}.New(raw)
|
||||
require.Error(t, err, "per-record tables validate too")
|
||||
}
|
||||
|
||||
// TestFactory_NilPricingSkipsEverything is the version-skew contract: a
|
||||
// new proxy under an old management ({} config) must build, allow, and
|
||||
// skip with unknown_model — degraded but never broken.
|
||||
func TestFactory_NilPricingSkipsEverything(t *testing.T) {
|
||||
mw := buildMiddleware(t, []byte("{}"))
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, middleware.DecisionAllow, out.Decision, "cost_meter always allows")
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
require.True(t, ok, "no pricing table means every request skips")
|
||||
assert.Equal(t, skipUnknownModel, value)
|
||||
}
|
||||
|
||||
func TestFactory_ConfigDefaultsPriceRequests(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -116,34 +170,78 @@ func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) {
|
||||
assert.Equal(t, "0.000750000", value, "0.00015 + 0.0006 per 1k tokens, 9-decimal format")
|
||||
}
|
||||
|
||||
func TestFactory_PricingPathOverride(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "custom.yaml"), []byte(fixturePricing), 0o600), "write custom pricing")
|
||||
builtin.Configure(context.Background(), dir, nil, nil, nil)
|
||||
|
||||
raw, err := json.Marshal(Config{PricingPath: "custom.yaml"})
|
||||
// TestInvoke_PerRecordEntryBeatsDefaults: when llm_router resolved a
|
||||
// provider record whose operator pinned a price for the model, that price
|
||||
// wins over the surface default.
|
||||
func TestInvoke_PerRecordEntryBeatsDefaults(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}},
|
||||
},
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-azure": {"gpt-4o": {InputPer1K: 0.005, OutputPer1K: 0.02}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
|
||||
mw := buildMiddleware(t, raw)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "2000"},
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: "prov-azure"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "cost.usd_total must be emitted with custom pricing path")
|
||||
assert.Equal(t, "0.015000000", value, "2*0.0025 + 1*0.01 = 0.015 with 9-decimal format")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0.025000000", value, "operator's per-record price (0.005+0.02) wins over the default (0.0025+0.01)")
|
||||
}
|
||||
|
||||
func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
// TestInvoke_PerRecordMissFallsBackToDefaults: a resolved record with no
|
||||
// entry for this model (or no entries at all) falls through to the
|
||||
// surface defaults — gateway providers rely on exactly this.
|
||||
func TestInvoke_PerRecordMissFallsBackToDefaults(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Defaults: map[string]map[string]pricing.EntryJSON{
|
||||
"openai": {"gpt-4o": {InputPer1K: 0.0025, OutputPer1K: 0.01}},
|
||||
},
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-1": {"some-other-model": {InputPer1K: 1, OutputPer1K: 1}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw := buildMiddleware(t, raw)
|
||||
|
||||
for name, recordID := range map[string]string{
|
||||
"record with other models": "prov-1",
|
||||
"record with no entries": "prov-gateway",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "openai"},
|
||||
{Key: middleware.KeyLLMModel, Value: "gpt-4o"},
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: recordID},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "1000"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "1000"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "per-record miss must fall back to the surface default, not skip")
|
||||
assert.Equal(t, "0.012500000", value, "default rates apply")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInvoke_NoResolvedProviderIDUsesDefaults: metadata without a
|
||||
// resolved provider id (router denied, or a chain without llm_router)
|
||||
// prices from the defaults table directly.
|
||||
func TestInvoke_NoResolvedProviderIDUsesDefaults(t *testing.T) {
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
|
||||
@@ -153,17 +251,15 @@ func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok, "cost.usd_total must be emitted")
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0.018000000", value, "0.003 + 0.015 = 0.018 with 9-decimal format")
|
||||
_, skipped := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
|
||||
assert.False(t, skipped, "cost.skipped must not be set when cost is computed")
|
||||
}
|
||||
|
||||
func TestInvoke_MissingProvider(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -179,8 +275,7 @@ func TestInvoke_MissingProvider(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_MissingModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -196,8 +291,7 @@ func TestInvoke_MissingModel(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_MissingTokens(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
@@ -240,8 +334,7 @@ func TestInvoke_MissingTokens(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_UnparseableTokens(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
@@ -271,8 +364,7 @@ func TestInvoke_UnparseableTokens(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_ZeroTokens(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -291,8 +383,7 @@ func TestInvoke_ZeroTokens(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_UnknownModel(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -309,8 +400,7 @@ func TestInvoke_UnknownModel(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvoke_NilInput(t *testing.T) {
|
||||
configureBuiltin(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfig(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
@@ -319,36 +409,12 @@ func TestInvoke_NilInput(t *testing.T) {
|
||||
assert.Empty(t, out.Metadata, "no metadata must be emitted on nil input")
|
||||
}
|
||||
|
||||
const fixturePricingWithCache = `openai:
|
||||
gpt-4o:
|
||||
input_per_1k: 0.0025
|
||||
output_per_1k: 0.01
|
||||
cached_input_per_1k: 0.00125
|
||||
anthropic:
|
||||
claude-sonnet-4-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
`
|
||||
|
||||
// configureBuiltinWithCacheRates points the package-level
|
||||
// FactoryContext at a tmp directory containing pricing entries that
|
||||
// include the cache rate fields.
|
||||
func configureBuiltinWithCacheRates(t *testing.T) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "pricing.yaml"), []byte(fixturePricingWithCache), 0o600), "write cache-aware pricing fixture")
|
||||
builtin.Configure(context.Background(), dir, nil, nil, nil)
|
||||
}
|
||||
|
||||
// TestInvoke_OpenAICachedSubsetDiscount proves the OpenAI shape end
|
||||
// to end through the middleware: cached_input_tokens is treated as a
|
||||
// SUBSET of input_tokens and discounted at the configured rate, not
|
||||
// added on top.
|
||||
func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) {
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -390,8 +456,7 @@ func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) {
|
||||
// shape: cache_read and cache_creation are additive to input_tokens
|
||||
// and each carries its own rate.
|
||||
func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) {
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -429,8 +494,37 @@ func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) {
|
||||
"output bucket bills 200 tokens at 0.015/1k")
|
||||
}
|
||||
|
||||
// TestInvoke_PerRecordEntryUsesSurfaceFormula: a per-record entry for an
|
||||
// anthropic-surface request must bill its cache buckets additively — the
|
||||
// formula follows llm.provider, not which table the entry came from.
|
||||
func TestInvoke_PerRecordEntryUsesSurfaceFormula(t *testing.T) {
|
||||
raw, err := json.Marshal(Config{Pricing: &PricingConfig{
|
||||
Providers: map[string]map[string]pricing.EntryJSON{
|
||||
"prov-ant": {"claude-sonnet-4-5": {InputPer1K: 0.003, OutputPer1K: 0.015, CacheReadPer1K: 0.0003, CacheCreationPer1K: 0.00375}},
|
||||
},
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
mw := buildMiddleware(t, raw)
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
{Key: middleware.KeyLLMProvider, Value: "anthropic"},
|
||||
{Key: middleware.KeyLLMModel, Value: "claude-sonnet-4-5"},
|
||||
{Key: middleware.KeyLLMResolvedProviderID, Value: "prov-ant"},
|
||||
{Key: middleware.KeyLLMInputTokens, Value: "256"},
|
||||
{Key: middleware.KeyLLMOutputTokens, Value: "200"},
|
||||
{Key: middleware.KeyLLMCachedInputTokens, Value: "768"},
|
||||
{Key: middleware.KeyLLMCacheCreationTokens, Value: "512"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, "0.005918400", value, "identical math to the defaults-table entry with the same rates")
|
||||
}
|
||||
|
||||
// assertBucket asserts one per-bucket cost key carries the expected
|
||||
// 6-decimal value.
|
||||
// 9-decimal value.
|
||||
func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) {
|
||||
t.Helper()
|
||||
got, ok := metaValue(t, md, key)
|
||||
@@ -439,13 +533,10 @@ func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) {
|
||||
}
|
||||
|
||||
// TestInvoke_CachedTokensAbsentFallsBackToBaseFormula covers the
|
||||
// "operator hasn't opted in" path: with no cached metadata keys
|
||||
// emitted, the meter must produce exactly the same cost as before
|
||||
// the feature landed. Critical so operators with the new binary but
|
||||
// no YAML changes see no behavioural drift on OpenAI requests.
|
||||
// no-cache-metadata path: with no cached keys emitted, the meter must
|
||||
// produce exactly the input+output cost.
|
||||
func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -460,7 +551,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
|
||||
require.True(t, ok)
|
||||
// 1000 input * 0.0025 + 500 output * 0.01 = 0.0025 + 0.005 = 0.0075
|
||||
assert.Equal(t, "0.007500000", value, "no cached metadata = same cost as before the feature landed")
|
||||
assert.Equal(t, "0.007500000", value, "no cached metadata = plain input+output cost")
|
||||
}
|
||||
|
||||
// TestInvoke_UnparseableCachedTokensSkippedSilently proves the
|
||||
@@ -469,8 +560,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
|
||||
// regular formula. Cache buckets are a refinement, never a reason to
|
||||
// abort cost computation.
|
||||
func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) {
|
||||
configureBuiltinWithCacheRates(t)
|
||||
mw := buildMiddleware(t, nil)
|
||||
mw := buildMiddleware(t, fixtureConfigWithCache(t))
|
||||
|
||||
out, err := mw.Invoke(context.Background(), &middleware.Input{
|
||||
Metadata: []middleware.KV{
|
||||
@@ -487,22 +577,10 @@ func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) {
|
||||
assert.Equal(t, "0.007500000", value, "same as the no-cached-metadata path")
|
||||
}
|
||||
|
||||
// TestMiddleware_CloseCancelsReloader proves Close stops the per-instance
|
||||
// pricing-reload goroutine: a chain rebuild retires the old instance and
|
||||
// calls Close, which must invoke the cancel func startReloader handed it so
|
||||
// the mtime-poll loop doesn't outlive the chain.
|
||||
func TestMiddleware_CloseCancelsReloader(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
m := newMiddleware(nil, cancel)
|
||||
|
||||
require.NoError(t, m.Close(), "Close must not error")
|
||||
require.Error(t, ctx.Err(), "Close must cancel the reloader context so the poll goroutine exits")
|
||||
}
|
||||
|
||||
// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) for an
|
||||
// instance with no reloader and for a nil receiver.
|
||||
// TestMiddleware_CloseNilSafe confirms Close is a no-op (no panic) even
|
||||
// for a nil receiver.
|
||||
func TestMiddleware_CloseNilSafe(t *testing.T) {
|
||||
require.NoError(t, newMiddleware(nil, nil).Close(), "no-reloader Close must be a no-op")
|
||||
require.NoError(t, newMiddleware(nil, nil).Close(), "Close must be a no-op")
|
||||
var m *Middleware
|
||||
require.NoError(t, m.Close(), "nil-receiver Close must be safe")
|
||||
}
|
||||
|
||||
@@ -10,11 +10,15 @@ import (
|
||||
)
|
||||
|
||||
// Config is the JSON-decoded shape accepted by the factory. The
|
||||
// runtime path consumes the normalised allowlist; raw config is not
|
||||
// runtime path consumes the normalised allowlists; raw config is not
|
||||
// retained beyond construction.
|
||||
type Config struct {
|
||||
ModelAllowlist []string `json:"model_allowlist"`
|
||||
PromptCapture PromptCapture `json:"prompt_capture"`
|
||||
// ProviderAllowlists maps a resolved provider id (KeyLLMResolvedProviderID) to
|
||||
// its model allowlist. A provider present is restricted to those models; one
|
||||
// absent is unrestricted. Kept per-provider so one provider's list can't leak
|
||||
// onto another.
|
||||
ProviderAllowlists map[string][]string `json:"provider_allowlists,omitempty"`
|
||||
PromptCapture PromptCapture `json:"prompt_capture"`
|
||||
}
|
||||
|
||||
// PromptCapture toggles the optional prompt capture + redaction step
|
||||
@@ -54,21 +58,28 @@ func isEmptyJSON(raw []byte) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// normaliseConfig lowercases and trims allowlist entries so the runtime
|
||||
// match is case-insensitive. Empty entries are dropped.
|
||||
// normaliseConfig lowercases and trims allowlist entries for case-insensitive
|
||||
// matching; empty entries drop. A provider whose entries all drop keeps an empty
|
||||
// (non-nil) list — "deny every model" — distinct from an absent provider
|
||||
// (unrestricted).
|
||||
func normaliseConfig(cfg Config) Config {
|
||||
if len(cfg.ModelAllowlist) == 0 {
|
||||
if len(cfg.ProviderAllowlists) == 0 {
|
||||
cfg.ProviderAllowlists = nil
|
||||
return cfg
|
||||
}
|
||||
cleaned := make([]string, 0, len(cfg.ModelAllowlist))
|
||||
for _, entry := range cfg.ModelAllowlist {
|
||||
n := normaliseModel(entry)
|
||||
if n == "" {
|
||||
continue
|
||||
cleaned := make(map[string][]string, len(cfg.ProviderAllowlists))
|
||||
for provider, models := range cfg.ProviderAllowlists {
|
||||
list := make([]string, 0, len(models))
|
||||
for _, entry := range models {
|
||||
n := normaliseModel(entry)
|
||||
if n == "" {
|
||||
continue
|
||||
}
|
||||
list = append(list, n)
|
||||
}
|
||||
cleaned = append(cleaned, n)
|
||||
cleaned[provider] = list
|
||||
}
|
||||
cfg.ModelAllowlist = cleaned
|
||||
cfg.ProviderAllowlists = cleaned
|
||||
return cfg
|
||||
}
|
||||
|
||||
|
||||
@@ -83,8 +83,9 @@ func (m *Middleware) MutationsSupported() bool { return false }
|
||||
// prompt capture only affects the metadata emitted alongside an allow.
|
||||
func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middleware.Output, error) {
|
||||
model, modelPresent := lookupMetadata(in.Metadata, middleware.KeyLLMModel)
|
||||
providerID, _ := lookupMetadata(in.Metadata, middleware.KeyLLMResolvedProviderID)
|
||||
|
||||
if denial := m.evaluateAllowlist(model, modelPresent); denial != nil {
|
||||
if denial := m.evaluateAllowlist(providerID, model, modelPresent); denial != nil {
|
||||
return denial, nil
|
||||
}
|
||||
|
||||
@@ -110,20 +111,32 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
|
||||
// is a no-op.
|
||||
func (m *Middleware) Close() error { return nil }
|
||||
|
||||
// evaluateAllowlist returns a deny Output when the configured allowlist
|
||||
// rejects the model. A nil return means the request should proceed.
|
||||
func (m *Middleware) evaluateAllowlist(model string, modelPresent bool) *middleware.Output {
|
||||
if len(m.cfg.ModelAllowlist) == 0 {
|
||||
// evaluateAllowlist denies when the resolved provider's allowlist rejects the
|
||||
// model; nil means proceed. Scoped to the provider llm_router resolved, so an
|
||||
// unrestricted provider (absent from config) is never caught by another's list.
|
||||
func (m *Middleware) evaluateAllowlist(providerID, model string, modelPresent bool) *middleware.Output {
|
||||
if len(m.cfg.ProviderAllowlists) == 0 {
|
||||
return nil
|
||||
}
|
||||
// Fail closed: with an allowlist configured, a request whose model the
|
||||
// upstream parser could not extract (absent or empty) must be denied rather
|
||||
// than allowed. This is what enforces the allowlist for URL/path-routed
|
||||
// providers (Bedrock, Vertex, ...) whose model lives outside the JSON body.
|
||||
// Restrictions exist but the resolved provider is unknown, so we can't tell
|
||||
// if this request targets a restricted provider — fail closed. llm_router
|
||||
// normally stamps the provider first, so this is a defensive guard.
|
||||
if providerID == "" {
|
||||
return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown)
|
||||
}
|
||||
allowlist, restricted := m.cfg.ProviderAllowlists[providerID]
|
||||
if !restricted {
|
||||
// This provider has no allowlist (some authorising policy left it
|
||||
// unrestricted); management owns any per-policy/group decision.
|
||||
return nil
|
||||
}
|
||||
// Fail closed: with an allowlist in effect for this provider, a request whose
|
||||
// model the parser couldn't extract (absent/empty) is denied. This enforces
|
||||
// the allowlist for path-routed providers (Bedrock, Vertex) with no body model.
|
||||
if !modelPresent || normaliseModel(model) == "" {
|
||||
return denyModel("", denyCodeModelUnknown, denyMessageModelUnknown, denyReasonModelUnknown)
|
||||
}
|
||||
if m.modelInAllowlist(model) {
|
||||
if modelInAllowlist(allowlist, model) {
|
||||
return nil
|
||||
}
|
||||
return denyModel(model, denyCodeModel, denyMessageModel, denyReasonModel)
|
||||
@@ -151,14 +164,15 @@ func denyModel(model, code, message, reason string) *middleware.Output {
|
||||
}
|
||||
}
|
||||
|
||||
// modelInAllowlist reports whether the model matches any allowlist
|
||||
// entry under the case-insensitive, trim-tolerant comparison rule.
|
||||
func (m *Middleware) modelInAllowlist(model string) bool {
|
||||
// modelInAllowlist reports whether the model matches any entry in the supplied
|
||||
// (already-normalised) allowlist under the case-insensitive, trim-tolerant
|
||||
// comparison rule.
|
||||
func modelInAllowlist(allowlist []string, model string) bool {
|
||||
normalised := normaliseModel(model)
|
||||
if normalised == "" {
|
||||
return false
|
||||
}
|
||||
for _, allowed := range m.cfg.ModelAllowlist {
|
||||
for _, allowed := range allowlist {
|
||||
if allowed == normalised {
|
||||
return true
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user